"""
Diagnostic tool to identify training issues and provide solutions.
"""

import os
import sys
import torch
import json
from pathlib import Path

# Windows encoding fix
if sys.platform == 'win32':
    if hasattr(sys.stdout, 'reconfigure'):
        sys.stdout.reconfigure(encoding='utf-8', errors='replace')

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

def check_system():
    """Check system configuration."""
    print("="*70)
    print("SYSTEM DIAGNOSTICS")
    print("="*70)
    
    issues = []
    suggestions = []
    
    # Python version
    python_version = sys.version_info
    print(f"\n[OK] Python: {python_version.major}.{python_version.minor}.{python_version.micro}")
    if python_version < (3, 7):
        issues.append("Python version should be >= 3.7")
        suggestions.append("Upgrade Python: https://www.python.org/downloads/")
    
    # PyTorch
    print(f"[OK] PyTorch: {torch.__version__}")
    
    # CUDA
    if torch.cuda.is_available():
        print(f"[OK] CUDA Available: Yes")
        print(f"  Device: {torch.cuda.get_device_name(0)}")
        print(f"  Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB")
        
        # Check CUDA version compatibility
        cuda_version = torch.version.cuda
        print(f"  CUDA Version: {cuda_version}")
    else:
        print(f"⚠ CUDA Available: No")
        issues.append("CUDA not available - training will be slow")
        suggestions.append(
            "Install CUDA-enabled PyTorch:\n"
            "  pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118"
        )
    
    return issues, suggestions


def check_data_files(data_dir='./data'):
    """Check data files."""
    print("\n" + "="*70)
    print("DATA FILES CHECK")
    print("="*70)
    
    issues = []
    suggestions = []
    
    required_files = ['sample_train.txt']
    optional_files = ['sample_valid.txt', 'sample_test.txt']
    
    if not os.path.exists(data_dir):
        issues.append(f"Data directory not found: {data_dir}")
        suggestions.append(f"Create data directory: mkdir {data_dir}")
        return issues, suggestions
    
    for filename in required_files:
        filepath = os.path.join(data_dir, filename)
        if not os.path.exists(filepath):
            issues.append(f"Required file missing: {filepath}")
            suggestions.append(
                f"Create sample data: python scripts/create_sample_data.py\n"
                f"Or download: python scripts/download_wikitext.py"
            )
        else:
            size = os.path.getsize(filepath)
            if size == 0:
                issues.append(f"File is empty: {filepath}")
                suggestions.append("Recreate data file")
            else:
                print(f"[OK] {filename}: {size / 1024:.1f} KB")
    
    for filename in optional_files:
        filepath = os.path.join(data_dir, filename)
        if os.path.exists(filepath):
            size = os.path.getsize(filepath)
            print(f"[OK] {filename}: {size / 1024:.1f} KB")
        else:
            print(f"[WARN] {filename}: Not found (optional)")
    
    return issues, suggestions


def check_checkpoints(checkpoint_dir='./checkpoints'):
    """Check checkpoint directory."""
    print("\n" + "="*70)
    print("CHECKPOINT DIRECTORY CHECK")
    print("="*70)
    
    issues = []
    suggestions = []
    
    if not os.path.exists(checkpoint_dir):
        print(f"[WARN] Directory doesn't exist: {checkpoint_dir}")
        print(f"  Will be created automatically")
    else:
        print(f"[OK] Directory exists: {checkpoint_dir}")
        
        # 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)
            print("[OK] Write permissions: OK")
        except Exception as e:
            issues.append(f"Cannot write to {checkpoint_dir}: {e}")
            suggestions.append(f"Check directory permissions or use a different location")
        
        # Check for existing checkpoints
        if os.path.isdir(checkpoint_dir):
            files = os.listdir(checkpoint_dir)
            checkpoint_files = [f for f in files if f.endswith('.pt')]
            if checkpoint_files:
                print(f"[OK] Found {len(checkpoint_files)} checkpoint file(s)")
                latest = max(checkpoint_files, key=lambda f: os.path.getmtime(os.path.join(checkpoint_dir, f)))
                print(f"  Latest: {latest}")
    
    return issues, suggestions


def check_training_logs(log_dir='./logs'):
    """Check for training logs and recent errors."""
    print("\n" + "="*70)
    print("TRAINING LOGS CHECK")
    print("="*70)
    
    issues = []
    suggestions = []
    
    if not os.path.exists(log_dir):
        print(f"[WARN] Log directory doesn't exist: {log_dir}")
        print("  Will be created automatically")
    else:
        print(f"[OK] Log directory exists: {log_dir}")
        log_files = [f for f in os.listdir(log_dir) if f.endswith('.log')]
        if log_files:
            print(f"[OK] Found {len(log_files)} log file(s)")
            
            # Check latest log for errors
            latest_log = max(log_files, key=lambda f: os.path.getmtime(os.path.join(log_dir, f)))
            log_path = os.path.join(log_dir, latest_log)
            
            try:
                with open(log_path, 'r', encoding='utf-8') as f:
                    lines = f.readlines()
                    error_lines = [l for l in lines[-50:] if 'ERROR' in l or 'CRITICAL' in l]
                    if error_lines:
                        print(f"[WARN] Found {len(error_lines)} error(s) in latest log:")
                        for line in error_lines[-5:]:  # Show last 5 errors
                            print(f"  {line.strip()[:100]}")
                        issues.append(f"Errors found in log: {latest_log}")
                        suggestions.append(f"Check log file: {log_path}")
            except Exception as e:
                print(f"⚠ Could not read log file: {e}")
    
    return issues, suggestions


def diagnose_checkpoint(checkpoint_path):
    """Diagnose a specific checkpoint."""
    print("\n" + "="*70)
    print(f"CHECKPOINT DIAGNOSIS: {checkpoint_path}")
    print("="*70)
    
    issues = []
    suggestions = []
    
    if not os.path.exists(checkpoint_path):
        issues.append(f"Checkpoint file not found: {checkpoint_path}")
        return issues, suggestions
    
    try:
        checkpoint = torch.load(checkpoint_path, map_location='cpu')
        
        print("[OK] Checkpoint loaded successfully")
        
        # Check required keys
        required_keys = ['model_state_dict', 'epoch']
        missing_keys = [k for k in required_keys if k not in checkpoint]
        if missing_keys:
            issues.append(f"Missing keys: {missing_keys}")
            suggestions.append("Checkpoint may be corrupted")
        else:
            print(f"[OK] Required keys present")
        
        # Display checkpoint info
        if 'epoch' in checkpoint:
            print(f"  Epoch: {checkpoint['epoch']}")
        if 'global_step' in checkpoint:
            print(f"  Global Step: {checkpoint['global_step']}")
        if 'best_val_loss' in checkpoint:
            print(f"  Best Val Loss: {checkpoint['best_val_loss']:.6f}")
        
        # Check model state dict
        if 'model_state_dict' in checkpoint:
            state_dict = checkpoint['model_state_dict']
            print(f"  Model parameters: {len(state_dict)} parameter groups")
            
            # Check for NaN in parameters
            has_nan = False
            for key, value in list(state_dict.items())[:10]:  # Check first 10
                if isinstance(value, torch.Tensor):
                    if torch.isnan(value).any():
                        has_nan = True
                        issues.append(f"NaN found in parameter: {key}")
            
            if not has_nan:
                print("[OK] No NaN detected in parameters (sampled)")
        
    except Exception as e:
        issues.append(f"Failed to load checkpoint: {e}")
        suggestions.append("Checkpoint may be corrupted or incompatible")
    
    return issues, suggestions


def main():
    """Run all diagnostics."""
    print("="*70)
    print("TRAINING DIAGNOSTIC TOOL")
    print("="*70)
    
    all_issues = []
    all_suggestions = []
    
    # Run all checks
    sys_issues, sys_suggestions = check_system()
    all_issues.extend(sys_issues)
    all_suggestions.extend(sys_suggestions)
    
    data_issues, data_suggestions = check_data_files()
    all_issues.extend(data_issues)
    all_suggestions.extend(data_suggestions)
    
    ckpt_issues, ckpt_suggestions = check_checkpoints()
    all_issues.extend(ckpt_issues)
    all_suggestions.extend(ckpt_suggestions)
    
    log_issues, log_suggestions = check_training_logs()
    all_issues.extend(log_issues)
    all_suggestions.extend(log_suggestions)
    
    # Check specific checkpoint if provided
    if len(sys.argv) > 1:
        checkpoint_path = sys.argv[1]
        ckpt_issues, ckpt_suggestions = diagnose_checkpoint(checkpoint_path)
        all_issues.extend(ckpt_issues)
        all_suggestions.extend(ckpt_suggestions)
    
    # Summary
    print("\n" + "="*70)
    print("DIAGNOSTIC SUMMARY")
    print("="*70)
    
    if not all_issues:
        print("\n[OK] All checks passed! System is ready for training.")
        return 0
    else:
        print(f"\n[WARN] Found {len(all_issues)} issue(s):")
        for i, issue in enumerate(all_issues, 1):
            print(f"  {i}. {issue}")
        
        if all_suggestions:
            print(f"\n[SUGGESTIONS]:")
            for i, suggestion in enumerate(all_suggestions, 1):
                print(f"  {i}. {suggestion}")
        
        return 1


if __name__ == '__main__':
    exit_code = main()
    sys.exit(exit_code)

