import os
import re
import random
import numpy as np
import pandas as pd
import soundfile as sf
from torch.utils import data
from suno_utils.audio import Audio
from suno_utils.utils.text import normalize_whitespace
from suno_utils.utils.lyrics import remove_speakers
from suno_utils.utils.text import read_jsonl


def clean_text(text: str) -> str:
    """General text cleaning. A bit tight but makes the content very clean.

    Returns a cleaned text string that is expected to be recongizable by hoot.
    """
    text = "\n" + text
    text = text.replace("’", "'").lower()
    text = text.replace('"', "").lower()
    text = re.sub(r"\[.+?\]", " ", text)  # tags
    text = re.sub(r"\n.+?\:", " ", text)  # new line ends with :
    text = re.sub(r"\n.+?\：", " ", text)  # new line ends with :
    text = re.sub(r"\n\(.+?\)", " ", text)  # new line with ()
    text = re.sub(r"[\d]", " ", text)  # digits
    text = re.sub(r"▁", "", text)  # special stuff
    text = remove_speakers(text)
    text = re.sub(r"[^\w\'\s]", " ", text)  # keep only the words
    text = normalize_whitespace(text)
    return text


class GeniusLyricDataset(data.Dataset):
    def __init__(
        self,
        data_path="/app/suno/data/audio_mono_24khz/genius_hq",
        split="train",
        input_length_s=30.0,
        num_samples=-1,
        is_concat=False,
    ):
        assert split in ["train", "valid"]
        self.data_path = data_path
        self.split = split
        self.input_length_s = input_length_s
        self.num_samples = num_samples
        self.fs = 24000
        self.is_concat = is_concat

        # get filelist
        if split == "train":
            self.ids = np.load(
                os.path.join(data_path, "metadata", "lyric_0.4", "train_ids.npy")
            )
            self.metadata = read_jsonl(
                os.path.join(data_path, "metadata", "lyric_0.4", "train_metadata.jsonl")
            )

        elif split == "valid":
            self.ids = np.load(
                os.path.join(data_path, "metadata", "lyric_0.4", "test_ids.npy")
            )
            self.metadata = read_jsonl(
                os.path.join(data_path, "metadata", "lyric_0.4", "test_metadata.jsonl")
            )
        print("There are %d audio files from Genius data" % len(self.ids))

    def __getitem__(self, index):
        # load audio
        track_id = self.ids[index]
        audio_fn = os.path.join(self.data_path, "audio", track_id + ".wav")
        audio = Audio.from_file(audio_fn, sample_rate=self.fs)
        wav = audio.array_float

        # get lyrics
        if self.split == "train":
            metadata = random.choice(self.metadata[index][1])

        elif self.split == "valid":
            metadata = self.metadata[index][1][0]

        lyrics = metadata["text"]
        start_s = metadata["start_s"]
        end_s = metadata["end_s"]
        start_ix = int(start_s * self.fs)
        end_ix = int(end_s * self.fs)

        # clearn lyrics
        lyrics = clean_text(lyrics)

        # append special tokens
        lyrics = "[CLS]" + "[Lyrics]" + lyrics

        # audio crop
        input_length = int(self.fs * self.input_length_s)
        wav = wav[start_ix:end_ix]
        if len(wav) < input_length:  # zero padding
            wav = np.pad(
                wav, (0, input_length - len(wav)), mode="constant", constant_values=0.0
            )
        wav = wav[:input_length]

        if self.is_concat:
            return wav.astype("float32"), lyrics
        else:
            return wav.astype("float32"), lyrics, track_id, start_s

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