"""
Simple tokenizer for text processing.
"""

import re
from typing import List, Dict, Optional
from collections import Counter


class SimpleTokenizer:
    """
    A simple character-level or word-level tokenizer.
    """
    
    def __init__(self, vocab: Optional[Dict[str, int]] = None, is_char_level: bool = False):
        self.is_char_level = is_char_level
        self.vocab = vocab or {}
        self.idx_to_token = {idx: token for token, idx in self.vocab.items()}
        
        # Special tokens
        self.pad_token = '<PAD>'
        self.unk_token = '<UNK>'
        self.bos_token = '<BOS>'
        self.eos_token = '<EOS>'
        
        # Ensure special tokens are in vocab
        if not self.vocab:
            self.vocab = {
                self.pad_token: 0,
                self.unk_token: 1,
                self.bos_token: 2,
                self.eos_token: 3,
            }
            self.idx_to_token = {0: self.pad_token, 1: self.unk_token, 2: self.bos_token, 3: self.eos_token}
    
    @property
    def pad_token_id(self):
        return self.vocab.get(self.pad_token, 0)
    
    @property
    def unk_token_id(self):
        return self.vocab.get(self.unk_token, 1)
    
    @property
    def vocab_size(self):
        return len(self.vocab)
    
    def build_vocab(self, texts: List[str], max_vocab_size: int = 10000, min_freq: int = 2):
        """
        Build vocabulary from texts.
        
        Args:
            texts: List of text strings
            max_vocab_size: Maximum vocabulary size
            min_freq: Minimum frequency for a token to be included
        """
        token_counter = Counter()
        
        for text in texts:
            if self.is_char_level:
                tokens = list(text)
            else:
                # Word-level: split on whitespace and punctuation
                tokens = re.findall(r'\w+|\S', text.lower())
            
            token_counter.update(tokens)
        
        # Create vocabulary
        self.vocab = {
            self.pad_token: 0,
            self.unk_token: 1,
            self.bos_token: 2,
            self.eos_token: 3,
        }
        
        # Add most common tokens
        idx = 4
        for token, count in token_counter.most_common(max_vocab_size - 4):
            if count >= min_freq:
                self.vocab[token] = idx
                idx += 1
        
        self.idx_to_token = {idx: token for token, idx in self.vocab.items()}
    
    def encode(self, text: str, add_special_tokens: bool = True) -> List[int]:
        """
        Encode text to token IDs.
        
        Args:
            text: Input text
            add_special_tokens: Whether to add BOS/EOS tokens
            
        Returns:
            List of token IDs
        """
        if self.is_char_level:
            tokens = list(text)
        else:
            tokens = re.findall(r'\w+|\S', text.lower())
        
        token_ids = []
        
        if add_special_tokens:
            token_ids.append(self.vocab[self.bos_token])
        
        for token in tokens:
            token_ids.append(self.vocab.get(token, self.vocab[self.unk_token]))
        
        if add_special_tokens:
            token_ids.append(self.vocab[self.eos_token])
        
        return token_ids
    
    def decode(self, token_ids: List[int], skip_special_tokens: bool = True) -> str:
        """
        Decode token IDs to text.
        
        Args:
            token_ids: List of token IDs
            skip_special_tokens: Whether to skip special tokens in output
            
        Returns:
            Decoded text string
        """
        tokens = []
        
        for token_id in token_ids:
            if token_id in self.idx_to_token:
                token = self.idx_to_token[token_id]
                if skip_special_tokens and token in [self.pad_token, self.bos_token, self.eos_token]:
                    continue
                if not skip_special_tokens or token != self.unk_token:
                    tokens.append(token)
        
        if self.is_char_level:
            return ''.join(tokens)
        else:
            return ' '.join(tokens)


# Convenience function for BPE-like tokenization (simplified)
def create_bpe_tokenizer(texts: List[str], vocab_size: int = 10000):
    """
    Create a simplified BPE-style tokenizer.
    This is a minimal implementation - for production, use sentencepiece or tiktoken.
    """
    return SimpleTokenizer(is_char_level=False)

