from data_gen import DataGenerator, Dataset, RawExample, TrivialAugmenter
from data import wavtool2midi, Vocab
import json
import torch
from data_tests import vocab_test_inputs
import pytest
from embedder import Embedder


class MockDataset(Dataset):
    def __init__(self, split_elems):
        super().__init__()
        self.split_elems = split_elems
        self.shuffle_called = False
        self.requested_ranksize = None
        self.requested_order = None

    def num_examples(self):
        return {split: len(elems) for split, elems in self.split_elems.items()}

    def stream_examples_impl(self, split, ranksize, order):
        self.requested_ranksize = ranksize
        self.requested_order = order
        elems = self.split_elems[split]
        for elem in elems:
            yield elem

    def shuffle(self):
        self.shuffle_called = True


def make_mock_dataset():
    def make_split(name):
        return [
            RawExample(
                id=i,
                descs=[f"split {name} example {i}"],
                example=wavtool2midi(
                    json.loads(v[0][0] if isinstance(v[0], list) else v[0])
                ),
                accomps=([] if v[1] is None else [wavtool2midi(json.loads(v[1]))]),
            )
            for i, v in enumerate(vocab_test_inputs)
        ]

    return MockDataset(
        {
            0: make_split("train"),
            1: make_split("eval"),
        }
    )


@pytest.mark.parametrize("batch_size", [1, 4, 64])
@pytest.mark.parametrize("forcebatches", [None, 2, 20])
@pytest.mark.parametrize("parallelism", [1, 4])
def test_data_generator_batches(batch_size, forcebatches, parallelism):
    really_test_data_generator(
        batch_size, 0, forcebatches, parallelism, "random", (0, 1)
    )


def test_data_generator_eval_split():
    really_test_data_generator(1, 1, None, 1, "random", (0, 1))


def test_data_generator_id_order():
    really_test_data_generator(1, 0, None, 1, "id", (0, 1))


def test_data_generator_ranksize():
    really_test_data_generator(1, 0, None, 1, "random", (3, 4))


@pytest.mark.parametrize("seq_len", [3, 7, 64])
@pytest.mark.parametrize("pack_batch", [False, True])
def test_data_generator_example_continuous(seq_len, pack_batch):
    dataset = make_mock_dataset()
    vocab = Vocab(24, 96, 24, 48, 20, 127, 2, 1, 24 * 4, 4, 4 * 4 * 24, 4, 24)

    data_gen = DataGenerator(
        dataset,
        vocab,
        0,
        (0, 1),
        1,
        seq_len,
        1,
        256,
        "tgt",
        1,
        "cpu",
        None,
        True,
        pack_batch,
        TrivialAugmenter(),
    )

    refs = []
    for raw_example in dataset.split_elems[0]:
        refs.append(
            vocab.midi_to_tensor(
                raw_example.example,
                accompany=(raw_example.accomps[0] if raw_example.accomps else None),
                strict=True,
            )
        )

    frags = []
    for batch in data_gen.generate(order="random"):
        s = batch.tgt[0]
        frags.append(s[s[:, 0] != vocab.pad.index])
    reconstructed = torch.cat(frags, dim=0)
    begins = torch.nonzero(reconstructed[:, 0] == vocab.begin.index).view(-1).tolist()
    begins.append(reconstructed.shape[0])
    for begin, end in zip(begins, begins[1:]):
        seq = reconstructed[begin:end]
        for i, ref in enumerate(refs):
            if torch.equal(seq, ref) or (
                ref[-1, 0] == vocab.end.index and torch.equal(seq, ref[:-1])
            ):
                del refs[i]
                break
        else:
            assert False, f"unexpected example in batch: {seq}"
    assert len(refs) == 0, "missing examples"


def really_test_data_generator(
    batch_size, split, forcebatches, parallelism, order, ranksize
):
    dataset = make_mock_dataset()
    vocab = Vocab(24, 96, 24, 48, 20, 127, 2, 1, 24 * 4, 4, 4 * 4 * 24, 4, 24)
    encoder = Embedder(load_pretrained_weights=False)

    data_gen = DataGenerator(
        dataset,
        vocab,
        split,
        ranksize,
        batch_size,
        128,
        1,
        None,
        "tgt",
        parallelism,
        "cpu",
        encoder.tokenize,
        False,
        False,
        TrivialAugmenter(),
    )
    assert not dataset.shuffle_called
    assert dataset.requested_ranksize is None
    for i in range(2):
        data_gen.shuffle()
        assert dataset.shuffle_called
        refs = []
        for raw_example in dataset.split_elems[split]:
            refs.append(
                (
                    vocab.midi_to_tensor(
                        raw_example.example,
                        accompany=(
                            raw_example.accomps[0] if raw_example.accomps else None
                        ),
                        strict=True,
                    ),
                    raw_example.descs,
                )
            )
        batch_count = 0
        for batch in data_gen.generate(force_batches=forcebatches, order=order):
            batch_count += 1
            is_empty_batch = True
            for i in range(batch.tgt.shape[0]):
                tgt_valid = batch.tgt[i][batch.tgt[i][:, 0] != vocab.pad.index]
                if tgt_valid.shape[0] == 0:
                    continue
                is_empty_batch = False
                for j, (ref_tensor, ref_descs) in enumerate(refs):
                    if (
                        torch.equal(tgt_valid, ref_tensor)
                        and ref_descs[0] == batch.texts[i]
                    ):
                        # verify encoding
                        enc_toks = batch.encoder_input_ids[
                            i, batch.encoder_attention_mask[i] > 0
                        ]
                        enc_str = encoder.detokenize(enc_toks.view(1, -1))[0]
                        assert enc_str == batch.texts[i]
                        del refs[j]
                        break
                else:
                    assert False, f"unexpected example in batch: {tgt_valid}"
        assert forcebatches is not None or not is_empty_batch, "unexpected empty batch"
        if forcebatches is None or forcebatches * batch_size >= len(
            dataset.split_elems[split]
        ):
            assert len(refs) == 0, f"missing examples from batches: {refs}"
        if forcebatches is not None:
            assert batch_count == forcebatches
        assert dataset.requested_order == order
        assert dataset.requested_ranksize == ranksize
