#!/usr/bin/env python
"""
Comprehensive import verification script for Suno GPT environment
Tests all critical packages and their functionality
"""

import sys
import importlib
import traceback
from typing import List, Tuple

def test_basic_imports() -> List[Tuple[str, str, bool, str]]:
    """Test basic package imports"""
    packages = [
        ("torch", "PyTorch"),
        ("torchvision", "TorchVision"),
        ("torchaudio", "TorchAudio"),
        ("transformers", "Transformers"),
        ("wandb", "Weights & Biases"),
        ("deepspeed", "DeepSpeed"),
        ("pytorch_lightning", "PyTorch Lightning"),
        ("nnAudio", "nnAudio"),
        ("auraloss", "Auraloss"),
        ("encodec", "Encodec"),
        ("flash_attn", "Flash Attention"),
        ("suno_utils", "Suno Utils"),
        ("g2p_en", "G2P English"),
        ("phonemizer", "Phonemizer"),
        ("sentencepiece", "SentencePiece"),
        ("tiktoken", "TikToken"),
        ("better_profanity", "Better Profanity"),
        ("einops", "Einops"),
        ("torchsde", "Torch SDE"),
        ("modal", "Modal"),
        ("ninja", "Ninja"),
        ("scipy", "SciPy"),
        ("numpy", "NumPy"),
        ("pandas", "Pandas"),
        ("tqdm", "TQDM"),
        ("joblib", "Joblib"),
        ("ffmpeg", "FFmpeg-Python"),
    ]
    
    results = []
    for package, name in packages:
        try:
            mod = importlib.import_module(package)
            version = getattr(mod, '__version__', 'unknown')
            results.append((name, version, True, "OK"))
        except ImportError as e:
            results.append((name, "N/A", False, str(e)[:50]))
        except Exception as e:
            results.append((name, "N/A", False, f"Error: {str(e)[:50]}"))
    
    return results

def test_cuda_functionality():
    """Test CUDA and GPU functionality"""
    print("\n" + "="*60)
    print("CUDA and GPU Tests")
    print("="*60)
    
    try:
        import torch
        
        print(f"PyTorch version: {torch.__version__}")
        print(f"CUDA available: {torch.cuda.is_available()}")
        
        if torch.cuda.is_available():
            print(f"CUDA version: {torch.version.cuda}")
            print(f"cuDNN version: {torch.backends.cudnn.version()}")
            print(f"Number of GPUs: {torch.cuda.device_count()}")
            
            for i in range(torch.cuda.device_count()):
                props = torch.cuda.get_device_properties(i)
                print(f"GPU {i}: {props.name}")
                print(f"  Memory: {props.total_memory / 1024**3:.1f} GB")
                print(f"  Compute Capability: {props.major}.{props.minor}")
            
            # Test basic CUDA operation
            try:
                x = torch.randn(100, 100).cuda()
                y = torch.randn(100, 100).cuda()
                z = torch.matmul(x, y)
                print("✓ CUDA tensor operations working")
            except Exception as e:
                print(f"✗ CUDA tensor operations failed: {e}")
        else:
            print("⚠ CUDA not available - CPU only mode")
            
    except ImportError:
        print("✗ PyTorch not available")
    except Exception as e:
        print(f"✗ Error testing CUDA: {e}")

def test_flash_attention():
    """Test Flash Attention functionality"""
    print("\n" + "="*60)
    print("Flash Attention Test")
    print("="*60)
    
    try:
        import torch
        import flash_attn
        from flash_attn import flash_attn_func
        
        print(f"Flash Attention version: {flash_attn.__version__}")
        
        # Test basic flash attention operation
        if torch.cuda.is_available():
            batch, heads, seq_len, dim = 2, 8, 128, 64
            q = torch.randn(batch, seq_len, heads, dim, device='cuda', dtype=torch.float16)
            k = torch.randn(batch, seq_len, heads, dim, device='cuda', dtype=torch.float16)
            v = torch.randn(batch, seq_len, heads, dim, device='cuda', dtype=torch.float16)
            
            try:
                out = flash_attn_func(q, k, v)
                print(f"✓ Flash Attention working - output shape: {out.shape}")
            except Exception as e:
                print(f"✗ Flash Attention operation failed: {e}")
        else:
            print("⚠ Skipping Flash Attention test - CUDA not available")
            
    except ImportError as e:
        print(f"✗ Flash Attention not available: {e}")
    except Exception as e:
        print(f"✗ Error testing Flash Attention: {e}")

def test_suno_utils():
    """Test suno_utils functionality"""
    print("\n" + "="*60)
    print("Suno Utils Test")
    print("="*60)
    
    try:
        import suno_utils
        print("✓ suno_utils imported successfully")
        
        # Try to import some submodules
        test_modules = [
            "suno_utils.audio",
            "suno_utils.io",
            "suno_utils.modeling",
        ]
        
        for module in test_modules:
            try:
                importlib.import_module(module)
                print(f"✓ {module} available")
            except ImportError:
                print(f"⚠ {module} not available")
                
    except ImportError as e:
        print(f"✗ suno_utils not available: {e}")
    except Exception as e:
        print(f"✗ Error testing suno_utils: {e}")

def test_audio_packages():
    """Test audio processing packages"""
    print("\n" + "="*60)
    print("Audio Package Tests")
    print("="*60)
    
    # Test nnAudio
    try:
        import nnAudio
        import torch
        from nnAudio.features import STFT
        
        if torch.cuda.is_available():
            stft = STFT(n_fft=2048, hop_length=512).cuda()
            test_audio = torch.randn(1, 16000).cuda()
            spec = stft(test_audio)
            print(f"✓ nnAudio STFT working - output shape: {spec.shape}")
        else:
            print("⚠ nnAudio test skipped - CUDA not available")
    except Exception as e:
        print(f"⚠ nnAudio test failed: {e}")
    
    # Test auraloss
    try:
        import auraloss
        import torch
        
        if torch.cuda.is_available():
            loss_fn = auraloss.freq.MultiResolutionSTFTLoss()
            x = torch.randn(1, 1, 16000).cuda()
            y = torch.randn(1, 1, 16000).cuda()
            loss = loss_fn(x, y)
            print(f"✓ Auraloss working - loss value: {loss.item():.4f}")
        else:
            print("⚠ Auraloss test skipped - CUDA not available")
    except Exception as e:
        print(f"⚠ Auraloss test failed: {e}")

def main():
    """Main verification function"""
    print("="*60)
    print("Suno GPT Environment Verification")
    print("="*60)
    print(f"Python: {sys.version}")
    print(f"Executable: {sys.executable}")
    print("="*60)
    
    # Test basic imports
    print("\nPackage Import Tests:")
    print("-"*40)
    
    results = test_basic_imports()
    
    # Print results in a nice table
    max_name_len = max(len(name) for name, _, _, _ in results)
    
    success_count = 0
    failed_packages = []
    
    for name, version, success, message in results:
        status = "✓" if success else "✗"
        if success:
            success_count += 1
            print(f"{status} {name:<{max_name_len}} {version:<15}")
        else:
            failed_packages.append(name)
            print(f"{status} {name:<{max_name_len}} {'FAILED':<15} {message}")
    
    print(f"\nSummary: {success_count}/{len(results)} packages imported successfully")
    
    if failed_packages:
        print(f"\n⚠ Failed packages: {', '.join(failed_packages)}")
    
    # Run functionality tests
    test_cuda_functionality()
    test_flash_attention()
    test_suno_utils()
    test_audio_packages()
    
    # Final summary
    print("\n" + "="*60)
    if not failed_packages:
        print("✓ All critical packages are available!")
        return 0
    else:
        print(f"⚠ Some packages failed to import. Please check the logs.")
        return 1

if __name__ == "__main__":
    sys.exit(main())