# %%
from suno_utils.audio import Audio
from suno_utils.utils.text import read_jsonl, read_json, write_jsonl
import random
import numpy as np
import matplotlib.pyplot as plt
from joblib import Parallel, delayed
from pathlib import Path
import time
import os
import pandas as pd


# %%
def calculate_spectral_features(spectrum_db, sample_rate=48000):
    """
    Calculate spectral features from magnitude spectrum in dB
    
    Returns:
    dict with spectral features in Hz
    """
    
    # Handle different input shapes
    if len(spectrum_db.shape) == 2:
        spectrum = spectrum_db[0]
    else:
        spectrum = spectrum_db
    
    # Convert from dB to linear magnitude
    magnitude = 10 ** (spectrum / 20.0)
    
    # Create frequency bins
    n_bins = len(spectrum)
    nyquist = sample_rate / 2
    frequencies = np.linspace(0, nyquist, n_bins)
    
    # Normalize magnitude
    total_magnitude = np.sum(magnitude)
    if total_magnitude == 0:
        return None
    
    normalized_mag = magnitude / total_magnitude
    
    # Calculate features
    centroid = np.sum(frequencies * normalized_mag)
    spread = np.sqrt(np.sum(((frequencies - centroid) ** 2) * normalized_mag))
    
    # Rolloff (85% energy cutoff)
    cumulative_mag = np.cumsum(normalized_mag)
    rolloff_idx = np.where(cumulative_mag >= 0.85)[0]
    rolloff = frequencies[rolloff_idx[0]] if len(rolloff_idx) > 0 else nyquist
    
    # Flatness (tonal vs noise)
    geometric_mean = np.exp(np.mean(np.log(magnitude[magnitude > 0])))
    arithmetic_mean = np.mean(magnitude)
    flatness = geometric_mean / arithmetic_mean if arithmetic_mean > 0 else 0
    
    return {
        'centroid': centroid,
        'spread': spread,
        'rolloff': rolloff,
        'flatness': flatness,
        'peak_freq': frequencies[np.argmax(magnitude)],
    }

# %%
def process_single_json(json_file_path, sample_rate=48000):
    """Process a single JSON file and return features"""
    
    try:
        data = read_json(json_file_path)
        
        # Extract average_spectrum_db
        if 'average_spectrum_db' not in data:
            return None
        
        features = {}
        if data.get('average_spectrum_db', None) is not None:
            spectrum_db = np.asarray(data['average_spectrum_db'])
            features = calculate_spectral_features(spectrum_db, sample_rate)

        mean_abs_stereo_diff = 0.0
        if data.get('average_stereo_spectrum_side', None) is not None:
            stereo_diff = np.asarray(data['average_stereo_spectrum_side'])
            mean_abs_stereo_diff = float(np.mean(np.abs(stereo_diff)))
        
        # Return filename and features
        return {
            'filename': Path(json_file_path).name,
            'centroid': features.get('centroid', None),
            'spread': features.get('spread', None),
            'rolloff': features.get('rolloff', None),
            'flatness': features.get('flatness', None),
            'mean_abs_stereo_diff': mean_abs_stereo_diff,
            'duration_s': data.get('duration_seconds', None),
            'lufs_db': data.get('lufs_db',None), 
            'lufs_db_factor': data.get('lufs_db_factor',None), 
            'rms_loudness_db': data.get('rms_loudness_db',None), 
            'peak_loudness_db': data.get('peak_loudness_db',None), 
            'stereo_width': data.get('stereo_width',None), 
            'clipped_samples': data.get('clipped_samples',None), 
        }
        
    except Exception as e:
        # Return error info for debugging if needed
        print(e)
        return None

def process_json_folder_parallel(folder_path, sample_rate=48000, n_jobs=20, batch_size=1000, progress_update=50000, max_files=None):
    """
    Process all JSON files in parallel with progress tracking
    
    Parameters:
    folder_path: Path to folder containing JSON files
    sample_rate: Audio sample rate in Hz
    n_jobs: Number of parallel jobs (-1 for all cores)
    batch_size: Number of files to process in each batch
    progress_update: Print progress every N files
    
    Returns:
    dict with lists of all features
    """
    
    folder = Path(folder_path)
    if not folder.exists():
        print(f"Folder {folder_path} does not exist!")
        return None
    
    # Get all JSON files
    json_files = list(folder.glob('*.json'))
    if max_files is not None:
        json_files = json_files[:max_files]
    total_files = len(json_files)
    
    print(f"Found {total_files:,} JSON files")
    print(f"Using {n_jobs if n_jobs > 0 else os.cpu_count()} parallel jobs")
    print(f"Processing in batches of {batch_size:,} files")
    
    # Initialize results
    all_results = []
    processed_count = 0
    start_time = time.time()
    
    # Process files in batches to manage memory
    for i in range(0, total_files, batch_size):
        batch_files = json_files[i:i + batch_size]
        batch_start = time.time()
        
        # Process batch in parallel
        batch_results = Parallel(n_jobs=n_jobs, backend='threading')(
            delayed(process_single_json)(json_file, sample_rate) 
            for json_file in batch_files
        )
        
        # Filter out None results and add to main results
        valid_results = [r for r in batch_results if r is not None]
        all_results.extend(valid_results)
        
        processed_count += len(batch_files)
        batch_time = time.time() - batch_start
        
        # Progress update
        if processed_count % progress_update == 0 or processed_count == total_files:
            elapsed_time = time.time() - start_time
            rate = processed_count / elapsed_time
            eta = (total_files - processed_count) / rate if rate > 0 else 0
            
            print(f"Processed {processed_count:,}/{total_files:,} files "
                  f"({processed_count/total_files*100:.1f}%) | "
                  f"Valid: {len(all_results):,} | "
                  f"Rate: {rate:.0f} files/sec | "
                  f"ETA: {eta/60:.1f} min | "
                  f"Batch time: {batch_time:.1f}s")
    
    # Convert to the expected format
    if not all_results:
        print("No valid results found!")
        return None
    
    total_time = time.time() - start_time
    print(f"\nCompleted! Processed {len(all_results):,} valid files in {total_time/60:.1f} minutes")
    print(f"Average rate: {len(all_results)/total_time:.0f} files/sec")
    
    return all_results

# %%
results = process_json_folder_parallel("/home/sara/sfx_analysis")

write_jsonl(results, "/home/sara/sfx_consolidated.jsonl")


