import os

os.environ["OPENBLAS_NUM_THREADS"] = "1"

import json
import numpy as np
import argparse
from tqdm.contrib.concurrent import process_map, thread_map
from multiprocessing import cpu_count
import librosa

from suno_utils.utils.text import read_jsonl
from suno_utils.audio import Audio

genius_data = read_jsonl("/app/suno/data/high_freq/genius_meta.jsonl")
print(len(genius_data))

n_worker = cpu_count() - 2
n_batch_size = 1000


def get_high_frequency_cutoff(job):
    try:
        job_id = list(job.keys())[0]
        job_file_path = job[job_id]
        audio = Audio.from_s3(job_file_path)
        # spectral roll-off
        roll_off = librosa.feature.spectral_rolloff(
            y=audio.array_float.astype("float32"),
            sr=48000,
            n_fft=4096,
            hop_length=480,
            roll_percent=0.99,
        )[0]
        # max pooling with 30s window
        window_step = 100
        pooled_roll_off = np.array(
            [
                np.max(roll_off[ix : ix + 3000])
                for ix in range(0, len(roll_off) - 3000 + window_step, window_step)
            ]
        )
        roll_off_median = np.median(pooled_roll_off)

        # get pass / fail
        # pass_fail = "pass" if roll_off_median > 17000 else "fail"
        return {job_id: roll_off_median}
    except:
        return {job_id: 0}


def get_high_frequency_cutoff_within_range(start_index=0, length=n_batch_size):
    output_json_path = f"/app/suno/data/high_freq/{start_index}.json"
    if os.path.exists(output_json_path):
        return
    # create a fake place holder
    with open(output_json_path, "w") as fp:
        json.dump({}, fp)
    output_result = {}
    jobs = []
    for index in range(start_index, min(start_index + length, len(genius_data))):
        jobs.append({genius_data[index]["id"]: genius_data[index]["audio_filepath"]})
    results = process_map(
        get_high_frequency_cutoff,
        jobs,
        max_workers=n_worker,
        chunksize=1,
    )
    for result in results:
        output_result.update(result)
    with open(output_json_path, "w") as fp:
        json.dump(output_result, fp)


def parse_args():
    parser = argparse.ArgumentParser()
    parser.add_argument("--start_index", type=int, default=0)
    parser.add_argument("--end_index", type=int, default=len(genius_data))
    args = parser.parse_args()
    return args


if __name__ == "__main__":
    input_args = parse_args()
    print(f"Start!!! {input_args.start_index, input_args.end_index}")
    for i in range(input_args.start_index, input_args.end_index, n_batch_size):
        get_high_frequency_cutoff_within_range(i, n_batch_size)
    print(f"DONE!!! {input_args.start_index, input_args.end_index}")
