import os
import json
import argparse
import tqdm
from tqdm.contrib.concurrent import process_map, thread_map
import numpy as np
import time
import datetime
import re
import json
import tqdm
import numpy as np
from suno_utils.audio import Audio
import math


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


if __name__ == "__main__":
    input_args = parse_args()
    with open("/home/tony/Data/Hoot/metas.json", "r") as fp:
        metas = json.load(fp)
    align_metas = []
    # will do this in batches of 100
    n_threads = 50  # A100 8, A10 16
    batch_size = 100
    for i in tqdm.tqdm(range(input_args.start_index, input_args.end_index, batch_size)):
        start_index = i
        end_index = min(start_index + batch_size, len(metas))
        if start_index > len(metas):
            continue
        output_path = (
            f"/home/tony/Data/Hoot/alignments/silence/batch_{start_index}.json"
        )
        if os.path.exists(output_path):
            continue
        # create an empyt file for place holder, in case there are racing conditions
        with open(output_path, "w") as fp:
            fp.write("")
        sub_metas = metas[start_index:end_index]

        def get_meta_output_based_on_id(test_id):
            # print('working on test_id', test_id)
            try:
                meta = metas[test_id]
                meta_id = meta["id"]
                audio_path = f'/app/suno/data/hoot/{meta["original_id"]}.wav'
                audio = Audio.from_file(audio_path)
                step_duration = 1
                silence_threshold = 0.02
                audio_loudness = [
                    max(
                        audio.get_segment(
                            from_s=i, to_s=min(i + step_duration, audio.duration_s)
                        ).loudness,
                        -1000,
                    )
                    for i in range(0, math.floor(audio.duration_s), step_duration)
                ]
                start_time = 0
                end_time = math.ceil(audio.duration_s)
                loundess_cut = np.quantile(audio_loudness, silence_threshold)
                while audio_loudness and audio_loudness[-1] < loundess_cut:
                    audio_loudness.pop()
                    end_time -= step_duration
                aligned_meta = {
                    "id": meta_id,
                    "duration_s": meta["duration_s"],
                    "end_duration": min(end_time, round(audio.duration_s, 1)),
                }
                return aligned_meta
            except Exception as e:
                print("oops", e)
                return None

        aligned_outputs = process_map(
            get_meta_output_based_on_id,
            list(range(start_index, end_index)),
            max_workers=n_threads,
            chunksize=1,
            disable=True,  # print or not
        )
        align_metas = [
            aligned_output
            for aligned_output in aligned_outputs
            if aligned_output is not None
        ]
        with open(output_path, "w") as fp:
            json.dump(align_metas, fp)
    print(f"DONE!!! {input_args.start_index, input_args.end_index}")
