import dataclasses
import re
from typing import Dict, List, Optional

import numpy as np


TEXT = "text"
TAG = "tag"
HESITATION = "--"
ANONYMOUS_SPEAKER_PTN = "speaker_{}"

ALLOWED_TRANSCRIPT_TAG = (HESITATION, "[laughter]")


def _all_or_none_are_none(*items):
    success = True
    first = items[0] is None
    for item in items[1:]:
        success = success and (first == (item is None))
    return success


@dataclasses.dataclass(frozen=True)
class Token:
    value: str
    type: str  # must be either "text" or "tag"
    speaker_id: Optional[str] = None
    start_s: Optional[float] = None
    end_s: Optional[float] = None
    metadata: Optional[Dict] = None

    def __post_init__(self):
        if self.type not in (TEXT, TAG):
            raise ValueError(f"bad token_type `{self.type}` found. Must be `{TEXT}` or `{TAG}`")
        if not _all_or_none_are_none(self.start_s, self.end_s):
            raise ValueError(
                f"start_s and end_s must both be float or none. Found {self.start_s}, {self.end_s}"
            )
        if self.type == TAG and self.value != "--" and (self.value[:1] != "[" or self.value[-1:] != "]"):
            raise ValueError(
                f"meta tags must be surrounded by brackets, with the exception of `--`, "
                f"but {self.value} was found"
            )

    @property
    def success(self):
        return self.start_s is not None and self.end_s is not None

    @property
    def exclude_in_transcript(self):
        return self.type == TAG and self.value not in ALLOWED_TRANSCRIPT_TAG

    def as_dict(self):
        """Convert to json serializable dict"""
        dict_rep = dataclasses.asdict(self)
        # remove None keys
        clean_dict_rep = {}
        for k, v in dict_rep.items():
            if v is not None:
                clean_dict_rep[k] = v
        return clean_dict_rep

    @classmethod
    def from_dict(cls, dict_rep):
        return cls(
            dict_rep["value"],
            dict_rep["type"],
            start_s=dict_rep.get("start_s"),
            end_s=dict_rep.get("end_s"),
            speaker_id=dict_rep.get("speaker_id"),
            metadata=dict_rep.get("metadata"),
        )

    def add_timestamps(self, start_s, end_s):
        return Token(
            self.value,
            self.type,
            start_s=start_s,
            end_s=end_s,
            speaker_id=self.speaker_id,
            metadata=self.metadata,
        )

    def add_speaker_id(self, speaker_id):
        return Token(
            self.value,
            self.type,
            start_s=self.start_s,
            end_s=self.end_s,
            speaker_id=speaker_id,
            metadata=self.metadata,
        )

    def add_metadata(self, metadata):
        return Token(
            self.value,
            self.type,
            start_s=self.start_s,
            end_s=self.end_s,
            speaker_id=self.speaker_id,
            metadata=metadata,
        )

    def remove_timestamps(self):
        return Token(self.value, self.type, speaker_id=self.speaker_id, metadata=self.metadata)

    def remove_speaker_id(self):
        return Token(
            self.value,
            self.type,
            start_s=self.start_s,
            end_s=self.end_s,
            metadata=self.metadata,
        )

    def remove_metadata(self):
        return Token(
            self.value,
            self.type,
            speaker_id=self.speaker_id,
            start_s=self.start_s,
            end_s=self.end_s,
        )

    # make alias methods
    to_dict = as_dict


@dataclasses.dataclass(frozen=True)
class Tokens:
    tokens: List[Token]

    @staticmethod
    def _verify_monotonic_timestamps(tokens):
        timestamps = []
        for t in tokens:
            if t.success:
                timestamps.append(t.start_s)
                timestamps.append(t.end_s)
        if len(timestamps) > 0 and np.min(np.diff(timestamps)) < 0:
            raise ValueError("token timestamps not monotonically increasing")

    def __post_init__(self):
        self._verify_monotonic_timestamps(self.tokens)

    def __len__(self):
        return len(self.tokens)

    def __getitem__(self, val):
        return self.tokens[val]

    def __repr__(self):
        repr_text = self.text[:10]
        if len(self.text) > 10:
            repr_text += "..."
        return f"Tokens(text=`{repr_text}`)"

    @property
    def success(self):
        text_tokens = [t for t in self.tokens if t.type == TEXT]
        return len(text_tokens) == 0 or any([t.success for t in text_tokens])

    @property
    def text(self):
        return " ".join(t.value for t in self.tokens)

    @property
    def plaintext(self):
        return " ".join(t.value for t in self.tokens if t.type == TEXT)

    @property
    def speaker_turns(self):
        turns = []
        cur_speaker = None
        tmp = []
        for token in self.tokens:
            if token.speaker_id is not None:
                if token.speaker_id != cur_speaker and len(tmp) > 0:
                    turns.append(
                        {
                            "speaker_id": cur_speaker,
                            "text": " ".join([t.value for t in tmp]),
                            "plaintext": Tokens(tmp).plaintext,
                            "tokens": [t.as_dict() for t in tmp],
                        }
                    )
                    tmp = []
                cur_speaker = token.speaker_id
            tmp.append(token)
        if len(tmp) > 0:
            turns.append(
                {
                    "speaker_id": cur_speaker,
                    "text": " ".join([t.value for t in tmp]),
                    "plaintext": Tokens(tmp).plaintext,
                    "tokens": [t.as_dict() for t in tmp],
                }
            )
        return turns

    @property
    def n_speakers(self):
        return len(set([t.speaker_id for t in self.tokens if t.speaker_id is not None]))

    @property
    def exclude_in_transcript(self):
        return any([t.type == TAG and t.value not in ALLOWED_TRANSCRIPT_TAG for t in self.tokens])

    def as_dict(self):
        """Convert to json serializable dict"""
        return [token.as_dict() for token in self.tokens]

    @classmethod
    def from_dict(cls, token_dicts):
        return cls([Token.from_dict(token_dict) for token_dict in token_dicts])

    @classmethod
    def from_text(cls, text, speaker_id=None):
        tokens = []
        for token_str in re.findall(r"\[.*?\]|[^\s]+", text):
            if re.match(r"\[.*\]", token_str) or token_str == HESITATION:
                token_type = TAG
            else:
                token_type = TEXT
            token = Token(token_str, token_type, speaker_id=speaker_id)
            tokens.append(token)
        return cls(tokens)

    @classmethod
    def from_speaker_turns(cls, turns):
        tokens = []
        for turn in turns:
            for token_str in re.findall(r"\[.*?\]|[^\s]+", turn["text"]):
                if re.match(r"\[.*\]", token_str):
                    token = Token(token_str, TAG)
                    tokens.append(token)
                elif token_str == HESITATION:
                    token = Token(token_str, TAG, speaker_id=turn["speaker_id"])
                    tokens.append(token)
                else:
                    token = Token(token_str, TEXT, speaker_id=turn["speaker_id"])
                    tokens.append(token)
        return cls(tokens)

    def anonymize_speakers(self):
        cur_speaker_n = 0
        speaker_id_map = {}
        anonymized_tokens = []
        for t in self.tokens:
            if t.speaker_id is None:
                anonymized_tokens.append(t)
                continue
            if t.speaker_id not in speaker_id_map:
                speaker_id_map[t.speaker_id] = cur_speaker_n
                cur_speaker_n += 1
            anonymized_tokens.append(t.add_speaker_id(speaker_id_map[t.speaker_id]))
        return Tokens(anonymized_tokens)

    def remove_timestamps(self):
        return Tokens([token.remove_timestamps() for token in self.tokens])

    def remove_speaker_ids(self):
        return Tokens([token.remove_speaker_id() for token in self.tokens])

    def remove_metadatas(self):
        return Tokens([token.remove_metadata() for token in self.tokens])

    # make alias methods
    to_dict = as_dict
