import os
import random
import json
import librosa
import numpy as np
import soundfile as sf
from torch.utils import data
from suno_utils.audio import Audio
from suno_utils.utils.text import read_jsonl


class ArtistYTMSDDataset(data.Dataset):
    def __init__(
        self,
        data_path="/app/suno/data/audio_mono_24khz/artist_ytmsd",
        split="train",
        input_length_s=30.0,
        sample_rate=24000,
        num_samples=-1,
    ):
        assert split in ["train", "valid"]
        self.data_path = data_path
        self.split = split
        self.input_length_s = input_length_s
        self.sample_rate = sample_rate
        self.num_samples = num_samples

        # load files
        self.metadata = read_jsonl(os.path.join(data_path, "%s.jsonl" % split))
        print("%d files are available for %s set" % (len(self.metadata), split))

    def concatenate_tags(self, tags):
        if self.split == "train":
            random.shuffle(tags)
        tags = [tag.lower() for tag in tags if len(tag.split(" ")) < 5]
        concatenated_tags = ", ".join(tags)

        # add [CLS] token
        concatenated_tags = "[CLS]" + concatenated_tags
        return concatenated_tags

    def __getitem__(self, index):
        # read data
        metadata = self.metadata[index]

        # sample two songs
        index_a, index_b = random.sample(metadata["indices"], 2)
        metadata_a = self.metadata[index_a]
        metadata_b = self.metadata[index_b]

        tags_a = metadata_a["clean_tags"]
        tags_b = metadata_b["clean_tags"]

        # load audio
        audio_path_a = metadata_a["filepath"]
        audio_path_b = metadata_b["filepath"]
        audio_a = Audio.from_file(audio_path_a, sample_rate=self.sample_rate)
        audio_b = Audio.from_file(audio_path_b, sample_rate=self.sample_rate)
        if self.split == "train":
            # random crop
            try:
                start_ms_a = random.randint(
                    0, (audio_a.duration_ms - int(1000 * self.input_length_s) - 1)
                )
                start_s_a = start_ms_a / 1000
                start_ms_b = random.randint(
                    0, (audio_b.duration_ms - int(1000 * self.input_length_s) - 1)
                )
                start_s_b = start_ms_b / 1000
            except ValueError as e:
                start_s_a = 0.0
                start_s_b = 0.0
            wav_a = audio_a.get_segment(
                from_s=start_s_a, to_s=start_s_a + self.input_length_s
            ).array_float
            wav_b = audio_b.get_segment(
                from_s=start_s_b, to_s=start_s_b + self.input_length_s
            ).array_float

        elif self.split == "valid":
            # crop first 30s
            wav_a = audio_a.get_segment(
                from_s=0.0, to_s=self.input_length_s
            ).array_float
            wav_b = audio_b.get_segment(
                from_s=0.0, to_s=self.input_length_s
            ).array_float

        if len(wav_a) < int(self.sample_rate * self.input_length_s):
            pad = int(self.sample_rate * self.input_length_s) - len(wav_a)
            wav_a = np.pad(wav_a, (0, pad), mode="constant", constant_values=0)
        if len(wav_b) < int(self.sample_rate * self.input_length_s):
            pad = int(self.sample_rate * self.input_length_s) - len(wav_b)
            wav_b = np.pad(wav_b, (0, pad), mode="constant", constant_values=0)

        # load tag labels
        concatenated_tags_a = self.concatenate_tags(tags_a)
        concatenated_tags_b = self.concatenate_tags(tags_b)

        return wav_a, wav_b, concatenated_tags_a, concatenated_tags_b

    def __len__(self):
        if self.num_samples > 0:
            return self.num_samples
        else:
            return len(self.metadata)
