"""
Quick script to check training progress.
"""

import os
import json
from datetime import datetime


def check_training_progress(checkpoint_dir='./checkpoints/benchmark_run'):
    """Check if training is running or completed."""
    
    errors = []
    warnings = []
    
    if not os.path.exists(checkpoint_dir):
        errors.append(f"Checkpoint directory not found: {checkpoint_dir}")
        errors.append("Training may not have started yet.")
        print("\n".join(errors))
        return
    
    print(f"Checking training progress in: {checkpoint_dir}")
    print("="*70)
    
    # Check for checkpoints
    checkpoints = [f for f in os.listdir(checkpoint_dir) if f.endswith('.pt')]
    
    if not checkpoints:
        print("No checkpoints found. Training may still be starting...")
        return
    
    print(f"\nFound {len(checkpoints)} checkpoint(s):")
    for checkpoint in sorted(checkpoints):
        path = os.path.join(checkpoint_dir, checkpoint)
        mtime = os.path.getmtime(path)
        size = os.path.getsize(path) / (1024 * 1024)  # MB
        time_str = datetime.fromtimestamp(mtime).strftime('%Y-%m-%d %H:%M:%S')
        print(f"  - {checkpoint}")
        print(f"    Size: {size:.1f} MB")
        print(f"    Modified: {time_str}")
    
    # Check for best model
    best_model = os.path.join(checkpoint_dir, 'best_model.pt')
    if os.path.exists(best_model):
        print(f"\n✓ Best model found: best_model.pt")
        mtime = os.path.getmtime(best_model)
        size = os.path.getsize(best_model) / (1024 * 1024)
        print(f"  Size: {size:.1f} MB")
        print(f"  Modified: {datetime.fromtimestamp(mtime).strftime('%Y-%m-%d %H:%M:%S')}")
    else:
        print(f"\n⚠ Best model not yet saved")
    
    # Check for evaluation results
    eval_results = [f for f in os.listdir(checkpoint_dir) if f.startswith('eval_results')]
    if eval_results:
        print(f"\nEvaluation results found: {len(eval_results)}")
        for result in sorted(eval_results):
            try:
                path = os.path.join(checkpoint_dir, result)
                with open(path, 'r') as f:
                    data = json.load(f)
                
                for dataset, metrics in data.items():
                    if dataset != 'model_info':
                        ppl = metrics.get('perplexity', 'N/A')
                        acc = metrics.get('accuracy', 'N/A')
                        print(f"  {dataset}: PPL={ppl:.2f if isinstance(ppl, float) else ppl}, Acc={acc*100:.2f if isinstance(acc, float) else acc}%")
            except:
                pass
    
    # Check for config
    config_path = os.path.join(checkpoint_dir, 'config.json')
    if os.path.exists(config_path):
        print(f"\n✓ Config file found")
        try:
            with open(config_path, 'r') as f:
                config = json.load(f)
            print(f"  Epochs: {config.get('training', {}).get('num_epochs', 'N/A')}")
            print(f"  Batch size: {config.get('training', {}).get('batch_size', 'N/A')}")
        except:
            pass
    
    print("\n" + "="*70)
    
    # Estimate if training is done
    if checkpoints:
        latest_checkpoint = max([os.path.join(checkpoint_dir, f) for f in checkpoints], 
                                key=os.path.getmtime)
        time_since_update = datetime.now().timestamp() - os.path.getmtime(latest_checkpoint)
        
        if time_since_update < 60:
            print("Status: [OK] Training appears to be active (recent checkpoint)")
        elif time_since_update < 300:
            print("Status: [WARN] Training may be paused or between epochs")
        else:
            print(f"Status: [WARN] Training appears to be completed or paused (no updates for {int(time_since_update/60)} minutes)")
        
        # Check for errors in latest checkpoint
        try:
            import torch
            latest = latest_checkpoint
            try:
                checkpoint = torch.load(latest, map_location='cpu')
                if 'best_val_loss' in checkpoint:
                    val_loss = checkpoint['best_val_loss']
                    if torch.isnan(torch.tensor(val_loss)) or torch.isinf(torch.tensor(val_loss)):
                        warnings.append(f"[WARN] Latest checkpoint has invalid validation loss")
            except Exception as e:
                warnings.append(f"[WARN] Could not read checkpoint {os.path.basename(latest)}: {e}")
                warnings.append("  -> Checkpoint may be corrupted")
        except Exception:
            pass  # If torch is not available, skip checkpoint validation
    else:
        print("Status: No checkpoints found yet")
    
    # Check logs for errors
    log_dir = os.path.join(os.path.dirname(checkpoint_dir), 'logs')
    if os.path.exists(log_dir):
        log_files = [f for f in os.listdir(log_dir) if f.endswith('.log')]
        if log_files:
            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', errors='ignore') as f:
                    lines = f.readlines()
                    error_lines = [l for l in lines[-100:] if 'ERROR' in l or 'CRITICAL' in l]
                    if error_lines:
                        print(f"\n[WARN] Found {len(error_lines)} error(s) in latest log: {latest_log}")
                        print("  Recent errors:")
                        for line in error_lines[-3:]:  # Show last 3
                            # Extract just the error message
                            msg = line.split(' - ', 2)[-1].strip()[:150]
                            print(f"    {msg}")
            except:
                pass
    
    if warnings:
        print("\n[WARN] Warnings:")
        for warning in warnings:
            print(f"  {warning}")


if __name__ == '__main__':
    import sys
    checkpoint_dir = sys.argv[1] if len(sys.argv) > 1 else './checkpoints/benchmark_run'
    check_training_progress(checkpoint_dir)


