"""
Data loading utilities.
"""

import os
import torch
from torch.utils.data import Dataset, DataLoader
from typing import List, Optional
from .tokenizer import SimpleTokenizer


class TextDataset(Dataset):
    """
    Dataset for text data.
    """
    
    def __init__(
        self,
        texts: List[str],
        tokenizer: SimpleTokenizer,
        max_length: int = 512,
        stride: int = None
    ):
        self.tokenizer = tokenizer
        self.max_length = max_length
        self.stride = stride or max_length // 2
        
        # Tokenize all texts
        self.tokenized_texts = []
        for text in texts:
            token_ids = tokenizer.encode(text, add_special_tokens=True)
            self.tokenized_texts.extend(token_ids)
        
        # Create sequences
        self.sequences = []
        if len(self.tokenized_texts) >= max_length:
            for i in range(0, len(self.tokenized_texts) - max_length + 1, self.stride):
                seq = self.tokenized_texts[i:i + max_length]
                if len(seq) == max_length:
                    self.sequences.append(seq)
        
        # If no sequences created (texts too short), create at least one padded sequence
        if len(self.sequences) == 0 and len(self.tokenized_texts) > 0:
            seq = self.tokenized_texts[:max_length]
            if len(seq) < max_length:
                seq = seq + [self.tokenizer.pad_token_id] * (max_length - len(seq))
            self.sequences.append(seq)
    
    def __len__(self):
        return len(self.sequences)
    
    def __getitem__(self, idx):
        sequence = self.sequences[idx]
        
        # Input and target (shifted by one)
        input_ids = torch.tensor(sequence[:-1], dtype=torch.long)
        target_ids = torch.tensor(sequence[1:], dtype=torch.long)
        
        # Create attention mask (all ones since we pad sequences to max_length)
        attention_mask = torch.ones_like(input_ids)
        
        return {
            'input_ids': input_ids,
            'target_ids': target_ids,
            'attention_mask': attention_mask
        }


def load_text_file(file_path: str) -> List[str]:
    """
    Load text from a file.
    
    Args:
        file_path: Path to text file
        
    Returns:
        List of text lines/paragraphs
    """
    with open(file_path, 'r', encoding='utf-8') as f:
        content = f.read()
    
    # Split into paragraphs or sentences
    paragraphs = [p.strip() for p in content.split('\n\n') if p.strip()]
    
    if not paragraphs:
        # Fall back to line-by-line
        paragraphs = [line.strip() for line in content.split('\n') if line.strip()]
    
    return paragraphs


def create_data_loaders(
    train_texts: List[str],
    val_texts: Optional[List[str]] = None,
    tokenizer: Optional[SimpleTokenizer] = None,
    max_length: int = 512,
    batch_size: int = 8,
    num_workers: int = 0,
    build_vocab: bool = True
) -> tuple:
    """
    Create data loaders for training.
    
    Args:
        train_texts: List of training texts
        val_texts: Optional list of validation texts
        tokenizer: Optional tokenizer (will be created if not provided)
        max_length: Maximum sequence length
        batch_size: Batch size
        num_workers: Number of dataloader workers
        build_vocab: Whether to build vocabulary from texts
        
    Returns:
        (train_loader, val_loader, tokenizer)
    """
    # Create tokenizer if needed
    if tokenizer is None:
        tokenizer = SimpleTokenizer(is_char_level=False)
        
        if build_vocab:
            tokenizer.build_vocab(train_texts + (val_texts or []))
    
    # Create datasets
    train_dataset = TextDataset(
        texts=train_texts,
        tokenizer=tokenizer,
        max_length=max_length
    )
    
    if len(train_dataset) == 0:
        raise ValueError(f"Training dataset is empty (0 sequences). Check input texts and max_length.")
    
    train_loader = DataLoader(
        train_dataset,
        batch_size=batch_size,
        shuffle=True,
        num_workers=num_workers,
        pin_memory=False  # Disable for Windows compatibility
    )
    
    val_loader = None
    if val_texts:
        val_dataset = TextDataset(
            texts=val_texts,
            tokenizer=tokenizer,
            max_length=max_length
        )
        
        val_loader = DataLoader(
            val_dataset,
            batch_size=batch_size,
            shuffle=False,
            num_workers=num_workers,
            pin_memory=True
        )
    
    return train_loader, val_loader, tokenizer

