import os
import sox
import pandas as pd
from tqdm import tqdm
from suno_utils.audio import Audio
from concurrent.futures import ProcessPoolExecutor

def pitch_shift_file(input_filepath: str, steps: int, output_filepath: str) -> bool:
    try:
        if not os.path.isfile(output_filepath):
            audio = Audio.from_file(input_filepath)
            sr = audio.sample_rate
            array_wav = audio.array_float
            
            tfm = sox.Transformer()
            tfm.pitch(steps)
            sox_shifted_wav = tfm.build_array(input_array=array_wav, sample_rate_in=sr)
            
            sox_audio = Audio.from_array_float(sox_shifted_wav, sr)
            Audio.write_mp3(sox_audio, output_filepath)
            
            return True
        else:
            print('File already exists.')
            return True
    except Exception as e:
        print(f"Error processing {input_filepath}: {e}")
        return False

def process_file(file):
    id = file[:-4]
    entry = tency_df.loc[id]
    orig_key = entry['key']
    
    results = []
    for shift in shifts:
        output_filepath = os.path.join(aug_data_dir, f'{id}_aug{shift}.mp3')
        success = pitch_shift_file(os.path.join(data_dir, file), shift, output_filepath)
        if success:
            results.append((id, shift, orig_key))
    
    return results

def update_progress(tqdm_obj, future):
    tqdm_obj.update(1)

if __name__ == '__main__':
    tency_df = pd.read_csv('tency_data.csv', index_col='ID')
    tency_df.index = tency_df.index.astype(str)

    data_dir = '/app/suno/christian_c/datasets/tency'
    aug_data_dir = '/app/suno/christian_c/datasets/aug_tency'
    shifts = [-5, -4, -3, -2, -1, 1, 2, 3, 4, 5]

    all_results = []
    num_files = len(os.listdir(data_dir))
    print(f'{num_files} available.')
    with tqdm(total=len(os.listdir(data_dir)), desc="Files processed") as pbar:
        with ProcessPoolExecutor() as executor:
            futures = [executor.submit(process_file, file) for file in os.listdir(data_dir)]
            for future in futures:
                future.add_done_callback(lambda p: update_progress(pbar, p))
                all_results.extend(future.result())

    major_labels = ['A major', 'Bb major', 'B major', 'C major', 'Db major',
                'D major', 'Eb major', 'E major', 'F major', 'F# major',
                'G major', 'Ab major']

    minor_labels = ['A minor', 'Bb minor', 'B minor',
                'C minor', 'C# minor', 'D minor', 'D# minor', 'E minor',
                'F minor', 'F# minor', 'G minor', 'G# minor']

    new_ids = []
    new_keys = []
    artists = []
    songs = []
    tempos = []

    for result in all_results:
        id, shift, orig_key = result
        entry = tency_df.loc[id]
        artist = entry['Artist']
        song = entry['Song']
        tempo = entry['tempo']
        
        artists.append(artist)
        songs.append(song)
        tempos.append(tempo)
        new_id = f'{id}_aug{shift}'
        new_ids.append(new_id)
        
        if orig_key in major_labels:
            new_key = major_labels[(major_labels.index(orig_key) + shift) % 12]
        elif orig_key in minor_labels:
            new_key = minor_labels[(minor_labels.index(orig_key) + shift) % 12]
        new_keys.append(new_key)

    aug_tency_df = pd.DataFrame(data={'ID': new_ids,
                                    'Artist': artists,
                                    'Song': songs,
                                    'tempo': tempos,
                                    'key': new_keys})
    aug_tency_df = aug_tency_df.set_index('ID')
    aug_tency_df.to_csv('aug_tency_data.csv')