import os
import random
import re

import numpy as np
from tokenizers import AddedToken
import torch
import torch.nn.functional as F
from transformers import PreTrainedTokenizerFast


# to avoid: "The current process just got forked, after parallelism has already been used"
os.environ["TOKENIZERS_PARALLELISM"] = "False"


global tokenizer
g_tokenizer = None


def _load_tokenizer(tokenizer_fp=None):
    global g_tokenizer
    if g_tokenizer is not None:
        return g_tokenizer
    # tokenizer = BertTokenizerFast.from_pretrained(
    #     "bert-base-multilingual-cased",
    #     model_max_length=512*4,
    # )
    assert os.path.exists(tokenizer_fp)
    g_tokenizer = PreTrainedTokenizerFast(
        tokenizer_file=tokenizer_fp,
        unk_token="[UNK]",
        pad_token="[PAD]",
    )
    g_tokenizer.add_special_tokens({"additional_special_tokens": [AddedToken("\n")]})
    return g_tokenizer


def _space_repl(m):
    s = m.group()
    n_newline = s.count("\n")
    if n_newline >= 2:
        return "\n\n"
    elif n_newline == 1:
        return "\n"
    return " "


def _simplify_whitespace(text, retain_newlines=True):
    """simplify while respecting up to 2 newlines"""
    if retain_newlines:
        text = re.sub(r"\s+", _space_repl, text).strip()
    else:
        text = re.sub(r"\s+", " ", text).strip()
    return text


def tokenize_batch(
    text_list,
    max_tokens=None,
    pad_token_id=0,
    retain_newlines=True,
    tokenizer_fp=None,
):
    tokenizer = _load_tokenizer(tokenizer_fp)
    text_list = [_simplify_whitespace(s, retain_newlines=retain_newlines) for s in text_list]
    text_enc = tokenizer(
        text_list,
        add_special_tokens=False,
        truncation=True,
        max_length=max_tokens,
        padding="longest",
        return_tensors="pt",
    )["input_ids"].type(torch.long)
    text_enc[text_enc == tokenizer.pad_token_id] = pad_token_id
    return text_enc


def _clean_tag(tag):
    return re.sub(r"\s+", " ", tag).strip()


def _augment_tag(s):
    if random.random() >= 0.95:
        s = s.upper()
    elif random.random() >= 0.95:
        s = s.capitalize()
    elif random.random() >= 0.9:
        s = s.title()
    elif random.random() >= 0.9:
        s = s.lower()
    if random.random() >= 0.5:
        s = s.replace("-", " ").strip()
    return s


# Structure:
# {start;vocals:start}  # is song start & vocals start within 8s of actual start
## tags go here
## lyrics go here
# {start;vocals:end}  # is song end & vocals end within 8s of actual end


def _get_start_control_tags(data_meta):
    start_s = data_meta.get("start_s")
    vocal_start_s = data_meta.get("vocal_start_s")
    control_tags = []
    if start_s is not None and start_s <= 0.5:
        control_tags.append("start")
        if vocal_start_s is not None and vocal_start_s <= 8:
            control_tags.append("vocals:start")
    if len(control_tags) == 0:
        return None
    return "{" + ";".join(control_tags) + "}"


def _get_end_control_tags(data_meta):
    end_s = data_meta.get("end_s")
    vocal_end_s = data_meta.get("vocal_end_s")
    original_duration_s = data_meta.get("original_duration_s")
    control_tags = []
    if end_s is not None and original_duration_s - end_s <= 0.5:
        control_tags.append("end")
        if vocal_end_s is not None and end_s - vocal_end_s <= 10:
            control_tags.append("vocals:end")
    if len(control_tags) == 0:
        return None
    return "{" + ";".join(control_tags) + "}"


def get_computed_tags(data_meta):
    computed_tags = []
    cutoff_freq = data_meta.get("cutoff_freq")
    if cutoff_freq is None:
        return computed_tags
    if cutoff_freq <= 16_000:
        computed_tags.append("low rolloff")
    if cutoff_freq >= 18_000:
        computed_tags.append("high rolloff")
    return computed_tags


def shift_codebooks(cfg, data_row, array_width=None):
    if array_width is None:
        array_width = (
            data_row.shape[-1]
            + cfg.semantic_n_codebooks * cfg.semantic_shift_factor
            + (cfg.coarse_n_codebooks - 1) * cfg.coarse_shift_factor
        )
    # build semantic
    y_semantic_arr = np.full(
        (cfg.semantic_n_codebooks, array_width), cfg.semantic_pad_token, dtype=np.int64
    )
    for n in range(cfg.semantic_n_codebooks):
        offs = n * cfg.semantic_shift_factor
        y_semantic_arr[n, offs : offs + data_row[n].shape[-1]] = data_row[n]
    # build coarse
    y_coarse_arr = np.full((cfg.coarse_n_codebooks, array_width), cfg.coarse_pad_token, dtype=np.int64)
    for n in range(cfg.coarse_n_codebooks):
        offs = cfg.semantic_n_codebooks * cfg.semantic_shift_factor + n * cfg.coarse_shift_factor
        n2 = cfg.semantic_n_codebooks + n
        y_coarse_arr[n, offs : offs + data_row[n2].shape[-1]] = data_row[n2]
    # combine audio
    y_audio_arr = np.concatenate([y_semantic_arr, y_coarse_arr], axis=0)
    return y_audio_arr


def get_sample(
    data_sampling_info,
    split,
    dataset_idx=None,
    rel_row_idx=None,
    use_private=False,
    inference=False,
    suppress_text=False,
    dummy_data=False,
    return_idx=False,
):
    if dummy_data:
        cfg = data_sampling_info["cfg"]
        x_audio_arr = np.zeros(
            (
                cfg.semantic_n_codebooks + cfg.coarse_n_codebooks,
                cfg.block_size - cfg.t_text,
            ),
            dtype=np.int64,
        )
        y_audio_arr = np.zeros(
            (
                cfg.semantic_n_codebooks + cfg.coarse_n_codebooks,
                cfg.block_size - cfg.t_text,
            ),
            dtype=np.int64,
        )
        return "", x_audio_arr, y_audio_arr
    data = data_sampling_info[split]["data"]
    metas = data_sampling_info[split]["metas"]
    if rel_row_idx is None:
        if dataset_idx is None:
            weights = data_sampling_info[split]["weights"]
            dataset_idx = random.choices(list(range(len(weights))), weights=weights, k=1)[0]
        idx_lists = data_sampling_info[split]["idx_lists"]
        rel_row_idx = random.choice(list(range(len(idx_lists[dataset_idx]))))
    rel_row_idx = rel_row_idx % len(idx_lists[dataset_idx])
    row_idx = idx_lists[dataset_idx][rel_row_idx]
    # names = data_sampling_info[split]["names"]
    # dataset_name = names[dataset_idx]
    data_row = data[row_idx].astype(np.int64)
    data_meta = metas[row_idx]
    cfg = data_sampling_info["cfg"]
    # TODO: change the transpose here
    data_row = data_row.T

    # split into chunks here
    n_tokens = np.where(data_row[0] != cfg.semantic_pad_token)[0][-1] + 1

    # TODO: this is a hack to not have to increase blocksize
    n_tokens = min(n_tokens, 3008 - 250)

    # abc -> acb
    # - predict only b
    # - a can be 0

    # take duration_s for calculation
    # 25% chance a is 0 (pre-painting)
    # max = min(duration_s-0.1, 59.9)
    # then pick a 0.01-max
    # then pick c 0.01-max
    # b is duration_s - a - c
    assert n_tokens >= 3
    c_len = random.randrange(1, n_tokens - 1)
    if random.random() <= -1:
        a_len = 0
    else:
        a_len = random.randrange(1, n_tokens - c_len)
    b_len = n_tokens - c_len - a_len
    assert a_len >= 0 and b_len > 0 and c_len > 0

    a_arr = data_row[:, :a_len]
    b_arr = data_row[:, a_len : a_len + b_len]
    c_arr = data_row[:, a_len + b_len : a_len + b_len + c_len]

    a_y_audio_arr = shift_codebooks(cfg, a_arr)
    b_y_audio_arr = shift_codebooks(cfg, b_arr)
    c_y_audio_arr = shift_codebooks(cfg, c_arr)

    a_len_post = a_y_audio_arr.shape[-1]
    b_len_post = b_y_audio_arr.shape[-1]
    c_len_post = c_y_audio_arr.shape[-1]

    assert a_len_post + c_len_post + b_len_post <= cfg.t_audio

    y_audio_arr = np.full(
        (cfg.semantic_n_codebooks + cfg.coarse_n_codebooks, cfg.t_audio),
        cfg.semantic_pad_token,
        dtype=np.int64,
    )
    y_audio_arr[cfg.semantic_n_codebooks :, :] = cfg.coarse_pad_token

    x_audio_arr = np.full(
        (cfg.semantic_n_codebooks + cfg.coarse_n_codebooks, cfg.t_audio),
        cfg.semantic_pad_token,
        dtype=np.int64,
    )
    x_audio_arr[cfg.semantic_n_codebooks :, :] = cfg.coarse_pad_token

    # assemble y
    y_audio_arr[:, :a_len_post] = a_y_audio_arr
    y_audio_arr[:, a_len_post + 1 : a_len_post + c_len_post + 1] = c_y_audio_arr
    y_audio_arr[:, a_len_post + c_len_post + 2 : a_len_post + c_len_post + b_len_post + 2] = (
        b_y_audio_arr
    )

    # make x and add infer
    x_audio_arr[: cfg.semantic_n_codebooks, 0:1] = cfg.semantic_infer_token
    x_audio_arr[cfg.semantic_n_codebooks :, 0:1] = cfg.coarse_infer_token
    x_audio_arr[:, 1 : 1 + a_len_post] = a_y_audio_arr
    x_audio_arr[: cfg.semantic_n_codebooks, 1 + a_len_post : 2 + a_len_post] = cfg.semantic_infer_token
    x_audio_arr[cfg.semantic_n_codebooks :, 1 + a_len_post : 2 + a_len_post] = cfg.coarse_infer_token
    x_audio_arr[:, 2 + a_len_post : 2 + a_len_post + c_len_post] = c_y_audio_arr
    x_audio_arr[
        : cfg.semantic_n_codebooks, 2 + a_len_post + c_len_post : 3 + a_len_post + c_len_post
    ] = cfg.semantic_infer_token
    x_audio_arr[
        cfg.semantic_n_codebooks :, 2 + a_len_post + c_len_post : 3 + a_len_post + c_len_post
    ] = cfg.coarse_infer_token
    x_audio_arr[:, 3 + a_len_post + c_len_post : 3 + a_len_post + c_len_post + b_len_post] = (
        b_y_audio_arr
    )

    # use text as indicator variable for abc
    # how to (optionally) encode duration
    # - fill in 100 tokens
    # - fill in arbitrary amount

    x_text_indicator = np.full((1, cfg.t_audio), cfg.text_pad_token, dtype=np.int64)
    x_text_indicator[:, : 1 + a_len_post] = cfg.text_pad_token + 1
    x_text_indicator[:, 1 + a_len_post : 2 + a_len_post + c_len_post] = cfg.text_pad_token + 3
    x_text_indicator[:, 2 + a_len_post + c_len_post : 3 + a_len_post + c_len_post + b_len_post] = (
        cfg.text_pad_token + 2
    )

    assert x_audio_arr.shape[-1] == cfg.block_size - cfg.t_text
    assert y_audio_arr.shape[-1] == cfg.block_size - cfg.t_text
    # build text
    text = ""
    # collect tags
    if use_private:
        tags = data_meta.get("tags_private", data_meta.get("tags", []))
    else:
        tags = data_meta.get("tags", [])
    # add computed tags
    computed_tags = get_computed_tags(data_meta)
    if len(computed_tags) > 0 and random.random() >= 0.1:
        tags.extend(computed_tags)
    # for tags remove newlines, empty tags, and case augment
    tags = [clean_tag for tag in tags if len(clean_tag := _clean_tag(tag)) > 0]
    if len(tags) > 0 and (inference or random.random() >= 0.25):
        if inference:
            text += f"[{', '.join(tags)[:128]}]\n\n"
        else:
            random.shuffle(tags)
            tags = tags[: random.randint(1, len(tags))]
            tags = [_augment_tag(tag) for tag in tags]
            tag_str = random.choice([", ", " ", "; "]).join(tags)
            text += f"[{tag_str[:128]}]"  # pretty arbitrary max len for now
            text += random.choice([" ", "\n", "\n\n"])
    # for tts_text sometimes remove metas, newlines and lower-case augment
    if use_private:
        tts_text = data_meta.get("text_private", data_meta.get("text", ""))
    else:
        tts_text = data_meta.get("text", "")
    if len(tts_text) > 0 and (inference or random.random() >= 0.1):
        if not inference:
            if random.random() >= 0.9:
                tts_text = tts_text.lower()
            if random.random() >= 0.9:
                tts_text = re.sub(r"\n+", " ", tts_text)
        text += tts_text.strip()
    # get control tags
    text = text.replace("{", "").replace("}", "")
    if not inference and random.random() >= 0.1:
        control_tags_start = _get_start_control_tags(data_meta)
        control_tags_end = _get_end_control_tags(data_meta)
        if control_tags_start is not None:
            text = control_tags_start + random.choice([" ", "\n", "\n\n"]) + text
        if control_tags_end is not None:
            text = text + random.choice([" ", "\n", "\n\n"]) + control_tags_end
    text = text.strip()
    if suppress_text:
        text = ""
    if return_idx:
        return row_idx, text, x_text_indicator, x_audio_arr, y_audio_arr
    return text, x_text_indicator, x_audio_arr, y_audio_arr


def get_batch(
    data_sampling_info,
    split,
    dataset_idx=None,
    row_idx=None,
    use_private=False,
    inference=False,
    min_text_offs=None,
    suppress_text=False,
    dummy_data=False,
    return_idx=False,
    n_offs=None,
):
    batch_size = data_sampling_info["batch_size"]
    device = data_sampling_info["device"]
    device_type = data_sampling_info["device_type"]
    tokenizer_fp = data_sampling_info.get("tokenizer_fp")
    cfg = data_sampling_info["cfg"]
    if not isinstance(dataset_idx, list):
        dataset_idx = [dataset_idx] * batch_size
    if not isinstance(row_idx, list):
        row_idx = [row_idx] * batch_size
    if n_offs is not None:
        row_idx = list(range(n_offs * batch_size, (n_offs + 1) * batch_size))
    x_text_list = []
    x_text_indicator_list = []
    x_audio_list = []
    y_list = []
    idx_list = []
    for n in range(batch_size):
        out = get_sample(
            data_sampling_info,
            split,
            dataset_idx=dataset_idx[n],
            rel_row_idx=row_idx[n],
            use_private=use_private,
            inference=inference,
            suppress_text=suppress_text,
            dummy_data=dummy_data,
            return_idx=return_idx,
        )
        if return_idx:
            idx, x_text, x_text_indicator, x_audio, y = out
            idx_list.append(idx)
        else:
            x_text, x_text_indicator, x_audio, y = out
        x_text_list.append(x_text)
        x_text_indicator_list.append(x_text_indicator)
        x_audio_list.append(torch.from_numpy(x_audio))
        y_list.append(torch.from_numpy(y))
    x_text = tokenize_batch(
        x_text_list,
        max_tokens=cfg.t_text,
        pad_token_id=cfg.text_pad_token,
        tokenizer_fp=tokenizer_fp,
    )
    if min_text_offs is not None and min_text_offs > x_text.shape[-1]:
        x_text = F.pad(
            x_text,
            (0, min_text_offs - x_text.shape[-1]),
            "constant",
            cfg.text_pad_token,
        )
    x_audio = torch.stack(x_audio_list)
    y = torch.stack(y_list)
    # combine all x and pad as much as needed
    x = torch.concatenate(
        [
            F.pad(
                x_audio[:, : cfg.semantic_n_codebooks],
                (
                    x_text.shape[-1],
                    cfg.block_size - x_text.shape[-1] - x_audio.shape[-1],
                ),
                "constant",
                cfg.semantic_pad_token,
            ),
            F.pad(
                x_audio[:, cfg.semantic_n_codebooks :],
                (
                    x_text.shape[-1],
                    cfg.block_size - x_text.shape[-1] - x_audio.shape[-1],
                ),
                "constant",
                cfg.coarse_pad_token,
            ),
        ],
        dim=1,
    )
    x = torch.concatenate(
        [
            F.pad(
                x_text[:, None],
                (0, cfg.block_size - x_text.shape[-1]),
                "constant",
                cfg.text_pad_token,
            ),
            x,
        ],
        dim=1,
    )
    text_offset = x_text.shape[-1]

    # add text indicator info
    for n, x_text_indicator in enumerate(x_text_indicator_list):
        x[n, :1, text_offset : text_offset + x_text_indicator.shape[-1]] = torch.from_numpy(
            x_text_indicator
        )
    assert x.shape == (
        batch_size,
        1 + cfg.semantic_n_codebooks + cfg.coarse_n_codebooks,
        cfg.block_size,
    )
    if device_type == "cuda":
        # pin arrays x,y, which allows us to move them to GPU asynchronously (non_blocking=True)
        x, y = (
            x.pin_memory().to(device, non_blocking=True),
            y.pin_memory().to(device, non_blocking=True),
        )
    else:
        x, y = x.to(device), y.to(device)
    del x_text_list, x_audio_list, y_list, x_text, x_audio
    if return_idx:
        return idx_list, text_offset, x, y
    return text_offset, x, y
