"""
Tokenizers for text processing.

Provides simple but complete tokenization for training and evaluation.
"""

import torch
from typing import List, Optional, Union, Dict
import re
from collections import Counter


class SimpleTokenizer:
    """
    Simple word-level tokenizer with special tokens.

    Good for testing and simple datasets. For production,
    consider using BPE or SentencePiece.
    """

    def __init__(
        self,
        vocab_size: int = 10000,
        min_frequency: int = 2,
        special_tokens: Optional[List[str]] = None
    ):
        self.vocab_size = vocab_size
        self.min_frequency = min_frequency

        # Special tokens
        if special_tokens is None:
            special_tokens = ['<pad>', '<unk>', '<bos>', '<eos>']
        self.special_tokens = special_tokens

        # Vocabularies
        self.token_to_id: Dict[str, int] = {}
        self.id_to_token: Dict[int, str] = {}

        # Special token IDs
        for idx, token in enumerate(special_tokens):
            self.token_to_id[token] = idx
            self.id_to_token[idx] = token

        self.pad_token_id = self.token_to_id.get('<pad>', 0)
        self.unk_token_id = self.token_to_id.get('<unk>', 1)
        self.bos_token_id = self.token_to_id.get('<bos>', 2)
        self.eos_token_id = self.token_to_id.get('<eos>', 3)

        self.is_trained = False

    def train(self, texts: List[str]):
        """
        Build vocabulary from texts.

        Args:
            texts: List of text strings
        """
        # Tokenize and count
        token_counts = Counter()

        for text in texts:
            tokens = self._tokenize_text(text)
            token_counts.update(tokens)

        # Filter by frequency and take top vocab_size
        filtered_tokens = [
            token for token, count in token_counts.items()
            if count >= self.min_frequency
        ]

        # Sort by frequency and take top vocab_size
        sorted_tokens = sorted(
            filtered_tokens,
            key=lambda t: token_counts[t],
            reverse=True
        )[:self.vocab_size - len(self.special_tokens)]

        # Build vocabulary
        next_id = len(self.special_tokens)
        for token in sorted_tokens:
            if token not in self.token_to_id:
                self.token_to_id[token] = next_id
                self.id_to_token[next_id] = token
                next_id += 1

        self.is_trained = True

    def _tokenize_text(self, text: str) -> List[str]:
        """
        Tokenize text into words.

        Simple whitespace + punctuation tokenization.
        """
        # Lowercase
        text = text.lower()

        # Add spaces around punctuation
        text = re.sub(r'([.,!?;:])', r' \1 ', text)

        # Split on whitespace
        tokens = text.split()

        return tokens

    def encode(
        self,
        text: str,
        add_special_tokens: bool = True,
        max_length: Optional[int] = None,
        padding: bool = False,
        truncation: bool = False
    ) -> List[int]:
        """
        Encode text to token IDs.

        Args:
            text: Input text
            add_special_tokens: Whether to add <bos> and <eos>
            max_length: Maximum sequence length
            padding: Whether to pad to max_length
            truncation: Whether to truncate to max_length

        Returns:
            token_ids: List of token IDs
        """
        if not self.is_trained:
            raise ValueError("Tokenizer not trained. Call train() first.")

        # Tokenize
        tokens = self._tokenize_text(text)

        # Convert to IDs
        token_ids = [
            self.token_to_id.get(token, self.unk_token_id)
            for token in tokens
        ]

        # Add special tokens
        if add_special_tokens:
            token_ids = [self.bos_token_id] + token_ids + [self.eos_token_id]

        # Truncate
        if truncation and max_length is not None and len(token_ids) > max_length:
            token_ids = token_ids[:max_length]

        # Pad
        if padding and max_length is not None:
            if len(token_ids) < max_length:
                token_ids = token_ids + [self.pad_token_id] * (max_length - len(token_ids))

        return token_ids

    def decode(
        self,
        token_ids: Union[List[int], torch.Tensor],
        skip_special_tokens: bool = True
    ) -> str:
        """
        Decode token IDs to text.

        Args:
            token_ids: Token IDs
            skip_special_tokens: Whether to skip special tokens

        Returns:
            text: Decoded text
        """
        if isinstance(token_ids, torch.Tensor):
            token_ids = token_ids.tolist()

        tokens = []
        for token_id in token_ids:
            token = self.id_to_token.get(token_id, '<unk>')

            if skip_special_tokens and token in self.special_tokens:
                continue

            tokens.append(token)

        text = ' '.join(tokens)
        return text

    def batch_encode(
        self,
        texts: List[str],
        add_special_tokens: bool = True,
        max_length: Optional[int] = None,
        padding: bool = True,
        truncation: bool = True,
        return_tensors: bool = True
    ) -> Union[List[List[int]], torch.Tensor]:
        """
        Encode batch of texts.

        Args:
            texts: List of texts
            add_special_tokens: Whether to add special tokens
            max_length: Maximum length
            padding: Whether to pad
            truncation: Whether to truncate
            return_tensors: Whether to return torch tensors

        Returns:
            token_ids: List of token ID lists or tensor [batch, seq_len]
        """
        encoded = [
            self.encode(
                text,
                add_special_tokens=add_special_tokens,
                max_length=max_length,
                padding=False,
                truncation=truncation
            )
            for text in texts
        ]

        # Pad to same length if needed
        if padding:
            if max_length is None:
                max_length = max(len(ids) for ids in encoded)

            encoded = [
                ids + [self.pad_token_id] * (max_length - len(ids))
                if len(ids) < max_length else ids
                for ids in encoded
            ]

        if return_tensors:
            return torch.tensor(encoded, dtype=torch.long)
        else:
            return encoded

    def save(self, path: str):
        """Save vocabulary to file."""
        import json

        data = {
            'vocab_size': self.vocab_size,
            'min_frequency': self.min_frequency,
            'special_tokens': self.special_tokens,
            'token_to_id': self.token_to_id,
            'id_to_token': {int(k): v for k, v in self.id_to_token.items()},
        }

        with open(path, 'w') as f:
            json.dump(data, f, indent=2)

    def load(self, path: str):
        """Load vocabulary from file."""
        import json

        with open(path, 'r') as f:
            data = json.load(f)

        self.vocab_size = data['vocab_size']
        self.min_frequency = data['min_frequency']
        self.special_tokens = data['special_tokens']
        self.token_to_id = data['token_to_id']
        self.id_to_token = {int(k): v for k, v in data['id_to_token'].items()}
        self.is_trained = True


class CharTokenizer:
    """
    Character-level tokenizer.

    Simpler than word-level, good for small datasets or languages
    without clear word boundaries.
    """

    def __init__(self, special_tokens: Optional[List[str]] = None):
        if special_tokens is None:
            special_tokens = ['<pad>', '<unk>', '<bos>', '<eos>']

        self.special_tokens = special_tokens
        self.char_to_id: Dict[str, int] = {}
        self.id_to_char: Dict[int, str] = {}

        # Add special tokens
        for idx, token in enumerate(special_tokens):
            self.char_to_id[token] = idx
            self.id_to_char[idx] = token

        self.pad_token_id = self.char_to_id.get('<pad>', 0)
        self.unk_token_id = self.char_to_id.get('<unk>', 1)
        self.bos_token_id = self.char_to_id.get('<bos>', 2)
        self.eos_token_id = self.char_to_id.get('<eos>', 3)

    def train(self, texts: List[str]):
        """Build character vocabulary."""
        chars = set()
        for text in texts:
            chars.update(text)

        # Sort for consistency
        sorted_chars = sorted(chars)

        next_id = len(self.special_tokens)
        for char in sorted_chars:
            if char not in self.char_to_id:
                self.char_to_id[char] = next_id
                self.id_to_char[next_id] = char
                next_id += 1

    def encode(self, text: str, add_special_tokens: bool = True) -> List[int]:
        """Encode text to character IDs."""
        char_ids = [
            self.char_to_id.get(char, self.unk_token_id)
            for char in text
        ]

        if add_special_tokens:
            char_ids = [self.bos_token_id] + char_ids + [self.eos_token_id]

        return char_ids

    def decode(self, char_ids: Union[List[int], torch.Tensor], skip_special_tokens: bool = True) -> str:
        """Decode character IDs to text."""
        if isinstance(char_ids, torch.Tensor):
            char_ids = char_ids.tolist()

        chars = []
        for char_id in char_ids:
            char = self.id_to_char.get(char_id, '<unk>')

            if skip_special_tokens and char in self.special_tokens:
                continue

            chars.append(char)

        return ''.join(chars)

    @property
    def vocab_size(self) -> int:
        """Get vocabulary size."""
        return len(self.char_to_id)
