#!/bin/bash

# Setup script following the original installation sequence exactly
# This script creates a conda environment with all Suno dependencies

set -e  # Exit on error

# Configuration
ENV_NAME="${1:-suno_env}"
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
LOG_FILE="${SCRIPT_DIR}/setup_${ENV_NAME}_$(date +%Y%m%d_%H%M%S).log"

# Colors for output
GREEN='\033[0;32m'
YELLOW='\033[1;33m'
RED='\033[0;31m'
BLUE='\033[0;34m'
NC='\033[0m' # No Color

# Logging functions
log() {
    echo -e "${GREEN}[$(date +'%Y-%m-%d %H:%M:%S')]${NC} $1" | tee -a "$LOG_FILE"
}

error() {
    echo -e "${RED}[ERROR]${NC} $1" | tee -a "$LOG_FILE"
    exit 1
}

warning() {
    echo -e "${YELLOW}[WARNING]${NC} $1" | tee -a "$LOG_FILE"
}

info() {
    echo -e "${BLUE}[INFO]${NC} $1" | tee -a "$LOG_FILE"
}

# Function to setup Suno Utils path
setup_suno_utils_path() {
    # Default path
    DEFAULT_SUNO_PATH="$HOME/projects/glockenspiel/suno_utils"
    
    # Ask user for Suno Utils path
    echo -e "${BLUE}Suno Utils Setup${NC}"
    echo "Enter the path to suno_utils package"
    echo "Press Enter to use default: $DEFAULT_SUNO_PATH"
    read -p "Path: " USER_SUNO_PATH
    
    # Use user path or default
    if [ -z "$USER_SUNO_PATH" ]; then
        SUNO_UTILS_PATH="$DEFAULT_SUNO_PATH"
        log "Using default suno_utils path: $SUNO_UTILS_PATH"
    else
        # Expand tilde if present
        SUNO_UTILS_PATH="${USER_SUNO_PATH/#\~/$HOME}"
        log "Using custom suno_utils path: $SUNO_UTILS_PATH"
    fi
    
    # Check if suno_utils exists
    if [ ! -d "$SUNO_UTILS_PATH" ]; then
        error "suno_utils not found at $SUNO_UTILS_PATH. Please provide a valid path to the suno_utils package."
    else
        log "suno_utils found at $SUNO_UTILS_PATH"
        # Check if it's a valid Python package
        if [ ! -f "$SUNO_UTILS_PATH/setup.py" ] && [ ! -f "$SUNO_UTILS_PATH/pyproject.toml" ]; then
            warning "$SUNO_UTILS_PATH doesn't appear to be a valid Python package"
            warning "Missing setup.py or pyproject.toml"
            read -p "Continue anyway? (y/n): " -n 1 -r
            echo
            if [[ ! $REPLY =~ ^[Yy]$ ]]; then
                error "Valid suno_utils package required"
            fi
        fi
    fi
}

# Function to setup Flash Attention path
setup_flash_attention_path() {
    # Default path
    DEFAULT_FLASH_PATH="$HOME/projects/flash-attention"
    
    # Ask user for Flash Attention path
    echo -e "${BLUE}Flash Attention Setup${NC}"
    echo "Enter the path to Flash Attention repository"
    echo "Press Enter to use default: $DEFAULT_FLASH_PATH"
    read -p "Path: " USER_FLASH_PATH
    
    # Use user path or default
    if [ -z "$USER_FLASH_PATH" ]; then
        FLASH_ATTN_PATH="$DEFAULT_FLASH_PATH"
        log "Using default Flash Attention path: $FLASH_ATTN_PATH"
    else
        # Expand tilde if present
        FLASH_ATTN_PATH="${USER_FLASH_PATH/#\~/$HOME}"
        log "Using custom Flash Attention path: $FLASH_ATTN_PATH"
    fi
    
    # Check if Flash Attention exists
    if [ ! -d "$FLASH_ATTN_PATH" ]; then
        warning "Flash Attention not found at $FLASH_ATTN_PATH"
        read -p "Would you like to clone Flash Attention to this location? (y/n): " -n 1 -r
        echo
        if [[ $REPLY =~ ^[Yy]$ ]]; then
            log "Cloning Flash Attention repository..."
            # Create parent directory if needed
            mkdir -p "$(dirname "$FLASH_ATTN_PATH")"
            git clone https://github.com/Dao-AILab/flash-attention.git "$FLASH_ATTN_PATH" >> "$LOG_FILE" 2>&1 || {
                error "Failed to clone Flash Attention repository"
            }
            
            # Initialize submodules
            cd "$FLASH_ATTN_PATH"
            log "Initializing Flash Attention submodules..."
            git submodule update --init --recursive >> "$LOG_FILE" 2>&1 || {
                warning "Failed to initialize some submodules, continuing anyway"
            }
            cd - > /dev/null
            
            log "Flash Attention cloned successfully!"
        else
            error "Flash Attention repository required. Please provide a valid path or allow cloning."
        fi
    else
        log "Flash Attention repository found at $FLASH_ATTN_PATH"
        # Update repository if it exists
        read -p "Would you like to update the Flash Attention repository? (y/n): " -n 1 -r
        echo
        if [[ $REPLY =~ ^[Yy]$ ]]; then
            log "Updating Flash Attention repository..."
            cd "$FLASH_ATTN_PATH"
            git pull >> "$LOG_FILE" 2>&1 || {
                warning "Failed to update repository, using existing version"
            }
            git submodule update --init --recursive >> "$LOG_FILE" 2>&1 || {
                warning "Failed to update submodules, continuing with existing"
            }
            cd - > /dev/null
        fi
    fi
}

# Main installation process
main() {
    log "Starting Suno environment setup for ${ENV_NAME}"
    log "Following original installation sequence"
    log "Log file: ${LOG_FILE}"
    
    # Setup Suno Utils path
    setup_suno_utils_path
    
    # Setup Flash Attention path
    setup_flash_attention_path
    
    # Step 1: List existing conda environments
    log "Step 1: Listing existing conda environments..."
    conda env list | tee -a "$LOG_FILE"
    
    # Check if environment already exists
    if conda env list | grep -q "^${ENV_NAME} "; then
        warning "Environment '${ENV_NAME}' already exists."
        read -p "Do you want to remove and recreate it? (y/n): " -n 1 -r
        echo
        if [[ $REPLY =~ ^[Yy]$ ]]; then
            log "Removing existing environment..."
            conda env remove -n "${ENV_NAME}" -y >> "$LOG_FILE" 2>&1
        else
            error "Environment already exists. Exiting."
        fi
    fi
    
    # Step 2: Create conda environment with Python 3.10
    log "Step 2: Creating conda environment with Python 3.10..."
    conda create -n "${ENV_NAME}" python=3.10 -y >> "$LOG_FILE" 2>&1
    log "Conda environment created successfully!"
    
    # Step 3: Activate environment and continue installation
    log "Step 3: Activating environment and installing packages..."
    
    # Source conda for this script
    source "$(conda info --base)/etc/profile.d/conda.sh"
    conda activate "${ENV_NAME}"
    
    # Step 4: Install suno_utils from glockenspiel
    log "Step 4: Installing suno_utils from ${SUNO_UTILS_PATH}..."
    cd "$SUNO_UTILS_PATH"
    pip install -e . >> "$LOG_FILE" 2>&1
    log "suno_utils installed successfully!"
    
    # Step 5: Uninstall torch (whatever version suno_utils installed)
    log "Step 5: Uninstalling torch installed by suno_utils..."
    # Log what version was installed
    EXISTING_TORCH=$(python -c "import torch; print(torch.__version__)" 2>/dev/null || echo "None")
    log "Current torch version: ${EXISTING_TORCH}"
    pip uninstall torch torchvision torchaudio -y >> "$LOG_FILE" 2>&1
    log "Torch uninstalled"
    
    # Step 6: Install PyTorch with CUDA 12.4 support
    log "Step 6: Installing PyTorch with CUDA 12.4 support..."
    pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu124 >> "$LOG_FILE" 2>&1
    
    # Verify PyTorch installation
    NEW_TORCH=$(python -c "import torch; print(torch.__version__)")
    CUDA_AVAILABLE=$(python -c "import torch; print(torch.cuda.is_available())")
    log "PyTorch ${NEW_TORCH} installed, CUDA available: ${CUDA_AVAILABLE}"
    
    # Step 7: Clone and install Flash Attention v2
    log "Step 7: Setting up Flash Attention v2..."
    
    cd "$FLASH_ATTN_PATH"
    
    # Install ninja
    log "Installing ninja..."
    pip install ninja >> "$LOG_FILE" 2>&1
    
    # Install Flash Attention v2
    log "Building and installing Flash Attention v2 (this may take 10-30 minutes)..."
    export MAX_JOBS=4  # Limit parallel jobs
    export FLASH_ATTENTION_FORCE_BUILD=TRUE
    python setup.py install >> "$LOG_FILE" 2>&1 || {
        warning "Flash Attention installation failed, continuing..."
    }
    
    # Step 8: Install additional Python packages
    log "Step 8: Installing additional Python packages..."
    pip install wandb nnAudio deepspeed auraloss torchsde g2p_en \
        transformers==4.44.0 modal==1.0.1 better_profanity >> "$LOG_FILE" 2>&1
    log "Additional packages installed!"
    
    # Step 9: Install sox via conda
    log "Step 9: Installing sox via conda..."
    conda install -c conda-forge sox -y >> "$LOG_FILE" 2>&1
    log "Sox installed!"
    
    # Return to script directory
    cd "${SCRIPT_DIR}"
    
    # Verification
    log "Running verification tests..."
    python -c "
import sys
packages = ['torch', 'transformers', 'wandb', 'deepspeed', 'nnAudio', 
            'auraloss', 'suno_utils', 'g2p_en', 'modal']
failed = []
for pkg in packages:
    try:
        __import__(pkg)
        print(f'✓ {pkg}')
    except ImportError as e:
        print(f'✗ {pkg}: {e}')
        failed.append(pkg)

import torch
print(f'PyTorch: {torch.__version__}')
print(f'CUDA available: {torch.cuda.is_available()}')

if failed:
    print(f'\\nWarning: Some packages failed to import: {failed}')
" | tee -a "$LOG_FILE"
    
    # Create activation script
    log "Creating activation script..."
    cat > "${SCRIPT_DIR}/activate_${ENV_NAME}.sh" << EOF
#!/bin/bash
# Activation script for ${ENV_NAME} environment

# Source conda
source "\$(conda info --base)/etc/profile.d/conda.sh"

# Activate environment
conda activate ${ENV_NAME}

# Set environment variables
export CUDA_VISIBLE_DEVICES=\${CUDA_VISIBLE_DEVICES:-0}
export OMP_NUM_THREADS=1
export TRITON_CACHE_DIR=/mnt/localdisk/.triton_cache_\$USER

echo "Environment ${ENV_NAME} activated!"
echo "Python: \$(which python)"
echo "PyTorch: \$(python -c 'import torch; print(torch.__version__)' 2>/dev/null || echo 'Not available')"
echo "CUDA available: \$(python -c 'import torch; print(torch.cuda.is_available())' 2>/dev/null || echo 'Unknown')"
echo "Flash Attention path: ${FLASH_ATTN_PATH}"
echo "Suno Utils path: ${SUNO_UTILS_PATH}"
EOF
    
    chmod +x "${SCRIPT_DIR}/activate_${ENV_NAME}.sh"
    
    log "================================================"
    log "Environment setup completed!"
    log "To activate the environment, run:"
    log "  source ${SCRIPT_DIR}/activate_${ENV_NAME}.sh"
    log "Or:"
    log "  conda activate ${ENV_NAME}"
    log "Flash Attention path: ${FLASH_ATTN_PATH}"
    log "Suno Utils path: ${SUNO_UTILS_PATH}"
    log "================================================"
}

# Run main function
main "$@"