"""
Pre-flight validation checks before training starts.
"""

import torch
import os
import sys
from typing import List, Tuple, Optional
from pathlib import Path


class TrainingValidator:
    """
    Validates system and configuration before training starts.
    Provides actionable error messages.
    """
    
    def __init__(self):
        self.errors = []
        self.warnings = []
        self.suggestions = []
    
    def validate_all(self, 
                    data_paths: dict,
                    model_config: dict,
                    training_config: dict,
                    checkpoint_dir: str = './checkpoints') -> Tuple[bool, List[str], List[str]]:
        """
        Run all validation checks.
        
        Returns:
            (is_valid, errors, warnings)
        """
        self.errors = []
        self.warnings = []
        self.suggestions = []
        
        # System checks
        self._check_cuda_availability()
        self._check_data_files(data_paths)
        self._check_model_config(model_config)
        self._check_training_config(training_config)
        self._check_checkpoint_dir(checkpoint_dir)
        self._check_gpu_memory(model_config, training_config)
        
        is_valid = len(self.errors) == 0
        
        return is_valid, self.errors, self.warnings
    
    def _check_cuda_availability(self):
        """Check CUDA availability and version."""
        if not torch.cuda.is_available():
            self.warnings.append(
                "CUDA not available. Training will run on CPU (much slower)."
            )
            self.suggestions.append(
                "Install CUDA-enabled PyTorch: pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118"
            )
        else:
            try:
                device_count = torch.cuda.device_count()
                device_name = torch.cuda.get_device_name(0)
                memory_gb = torch.cuda.get_device_properties(0).total_memory / 1e9
                
                self.suggestions.append(
                    f"GPU detected: {device_name} ({memory_gb:.1f} GB)"
                )
                
                if memory_gb < 4:
                    self.warnings.append(
                        f"GPU has only {memory_gb:.1f} GB. Consider reducing batch_size or model size."
                    )
            except Exception as e:
                self.errors.append(f"Error checking GPU: {e}")
    
    def _check_data_files(self, data_paths: dict):
        """Check that data files exist and are readable."""
        required = ['train']
        optional = ['val', 'test']
        
        for key in required:
            if key not in data_paths:
                self.errors.append(f"Required data file '{key}' not specified.")
                continue
            
            path = data_paths[key]
            if not os.path.exists(path):
                self.errors.append(
                    f"Training data file not found: {path}\n"
                    f"  → Create sample data: python scripts/create_sample_data.py\n"
                    f"  → Or download data: python scripts/download_wikitext.py"
                )
            elif os.path.getsize(path) == 0:
                self.errors.append(
                    f"Training data file is empty: {path}\n"
                    f"  → File exists but contains no data"
                )
            else:
                # Check if file is readable
                try:
                    with open(path, 'r', encoding='utf-8') as f:
                        first_line = f.readline()
                        if not first_line.strip():
                            self.warnings.append(f"Data file {path} appears empty or starts with empty line.")
                except UnicodeDecodeError:
                    self.errors.append(
                        f"Data file {path} is not valid UTF-8 text.\n"
                        f"  → Ensure file is plain text, not binary"
                    )
                except Exception as e:
                    self.errors.append(f"Cannot read data file {path}: {e}")
        
        for key in optional:
            if key in data_paths:
                path = data_paths[key]
                if not os.path.exists(path):
                    self.warnings.append(f"Optional {key} data file not found: {path}")
    
    def _check_model_config(self, config: dict):
        """Validate model configuration."""
        required_keys = ['vocab_size', 'embedding_dim', 'num_layers']
        
        for key in required_keys:
            if key not in config:
                self.errors.append(f"Model config missing required key: {key}")
        
        # Check value ranges
        if 'vocab_size' in config:
            vocab_size = config['vocab_size']
            if not isinstance(vocab_size, int) or vocab_size < 100:
                self.warnings.append(
                    f"Vocab size {vocab_size} is very small. Recommended: 1000-50000"
                )
            if vocab_size > 100000:
                self.warnings.append(
                    f"Vocab size {vocab_size} is very large. May cause memory issues."
                )
        
        if 'embedding_dim' in config:
            embed_dim = config['embedding_dim']
            if embed_dim % 2 != 0:
                self.warnings.append(
                    f"Embedding dimension {embed_dim} is not even. Some operations work better with even numbers."
                )
            if embed_dim < 64:
                self.warnings.append(f"Embedding dimension {embed_dim} is very small.")
        
        if 'num_layers' in config:
            num_layers = config['num_layers']
            if num_layers < 1:
                self.errors.append("Number of layers must be >= 1")
            if num_layers > 24:
                self.warnings.append(
                    f"Number of layers {num_layers} is very high. Training will be slow and memory-intensive."
                )
        
        if 'max_seq_length' in config:
            seq_len = config['max_seq_length']
            if seq_len > 2048:
                self.warnings.append(
                    f"Max sequence length {seq_len} is very large. May cause OOM errors."
                )
    
    def _check_training_config(self, config: dict):
        """Validate training configuration."""
        if 'batch_size' in config:
            batch_size = config['batch_size']
            if batch_size < 1:
                self.errors.append("Batch size must be >= 1")
            if batch_size > 64:
                self.warnings.append(
                    f"Batch size {batch_size} is large. May cause OOM errors."
                )
        
        if 'learning_rate' in config:
            lr = config['learning_rate']
            if lr <= 0:
                self.errors.append("Learning rate must be > 0")
            if lr > 1:
                self.warnings.append(f"Learning rate {lr} is very high. May cause training instability.")
            if lr < 1e-6:
                self.warnings.append(f"Learning rate {lr} is very low. Training may be very slow.")
        
        if 'num_epochs' in config:
            epochs = config['num_epochs']
            if epochs < 1:
                self.errors.append("Number of epochs must be >= 1")
    
    def _check_checkpoint_dir(self, checkpoint_dir: str):
        """Check checkpoint directory."""
        try:
            os.makedirs(checkpoint_dir, exist_ok=True)
            
            # Check write permissions
            test_file = os.path.join(checkpoint_dir, '.write_test')
            try:
                with open(test_file, 'w') as f:
                    f.write('test')
                os.remove(test_file)
            except Exception as e:
                self.errors.append(
                    f"Cannot write to checkpoint directory {checkpoint_dir}: {e}\n"
                    f"  → Check directory permissions or use a different location"
                )
        except Exception as e:
            self.errors.append(
                f"Cannot create checkpoint directory {checkpoint_dir}: {e}\n"
                f"  → Check parent directory permissions"
            )
    
    def _check_gpu_memory(self, model_config: dict, training_config: dict):
        """Estimate GPU memory requirements."""
        if not torch.cuda.is_available():
            return
        
        try:
            available_memory = torch.cuda.get_device_properties(0).total_memory / 1e9
            
            # Rough estimate (very approximate)
            embed_dim = model_config.get('embedding_dim', 512)
            vocab_size = model_config.get('vocab_size', 10000)
            num_layers = model_config.get('num_layers', 6)
            batch_size = training_config.get('batch_size', 8)
            seq_len = model_config.get('max_seq_length', 512)
            
            # Rough estimate: parameters * 4 bytes + activations
            # This is very approximate
            param_memory = (
                vocab_size * embed_dim * 2 +  # embeddings
                embed_dim * embed_dim * num_layers * 12 +  # transformer layers
                vocab_size * embed_dim  # output head
            ) * 4 / 1e9  # 4 bytes per float
            
            # Activation memory (very rough)
            activation_memory = batch_size * seq_len * embed_dim * num_layers * 8 * 4 / 1e9
            
            estimated_total = param_memory + activation_memory + 1  # +1GB overhead
            
            if estimated_total > available_memory:
                self.warnings.append(
                    f"Estimated memory requirement ({estimated_total:.1f} GB) exceeds "
                    f"available GPU memory ({available_memory:.1f} GB).\n"
                    f"  → Reduce batch_size (try {max(1, batch_size // 2)})\n"
                    f"  → Reduce max_seq_length (try {seq_len // 2})\n"
                    f"  → Reduce embedding_dim (try {embed_dim // 2})\n"
                    f"  → Reduce num_layers (try {max(1, num_layers - 2)})"
                )
        except Exception as e:
            self.warnings.append(f"Could not estimate GPU memory: {e}")
    
    def print_report(self):
        """Print validation report."""
        print("\n" + "="*70)
        print("PRE-FLIGHT VALIDATION")
        print("="*70)
        
        if self.suggestions:
            print("\n✓ System Information:")
            for suggestion in self.suggestions:
                print(f"  {suggestion}")
        
        if self.warnings:
            print("\n⚠ Warnings:")
            for warning in self.warnings:
                print(f"  {warning}")
        
        if self.errors:
            print("\n✗ Errors:")
            for error in self.errors:
                print(f"  {error}")
        else:
            print("\n✓ All validation checks passed!")
        
        print("="*70 + "\n")


def validate_training_setup(data_paths: dict, model_config: dict, 
                          training_config: dict, checkpoint_dir: str = './checkpoints') -> bool:
    """
    Quick validation function.
    
    Returns:
        True if all checks pass, False otherwise
    """
    validator = TrainingValidator()
    is_valid, errors, warnings = validator.validate_all(
        data_paths, model_config, training_config, checkpoint_dir
    )
    validator.print_report()
    
    return is_valid


