import os
import time
import glob
import torch
import argparse
import torchaudio
import numpy as np
import scipy.signal
import pyloudnorm as pyln
import multiprocessing as mp

from tqdm import tqdm

from suno_amp.v0 import apply_postprocessing

VERSIONS = [0]


def run(filepath: str, output_dir: str, version: int = 0):
    audio, sample_rate = torchaudio.load(filepath)
    filename = os.path.basename(filepath).split(".")[0]
    print(filename)

    # start = time.perf_counter()
    if version == 0:
        output_audio = apply_postprocessing(audio, sample_rate)
    else:
        raise ValueError(f"Invalid version: {version}. Must be one of {VERSIONS}.")
    # end = time.perf_counter()
    # elapsed = timings.append(end - start)
    # print(np.mean(elapsed))

    # loudnorm the input for comparision
    # loudness normalize to target LUFS dB
    ffmpeg_filter = f"loudnorm=I=-16.0"
    effector = torchaudio.io.AudioEffector(ffmpeg_filter)
    audio_norm = effector.apply(audio.T, sample_rate).T

    # save output
    out_filepath_input = os.path.join(output_dir, filename + ".wav")
    out_filepath_output = os.path.join(output_dir, filename + "-output.wav")
    torchaudio.save(out_filepath_input, audio_norm, sample_rate)
    torchaudio.save(out_filepath_output, output_audio, sample_rate)


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument(
        "input_dir", help="Path to directory containing audio files to process."
    )
    parser.add_argument("--version", default=0, type=int)
    args = parser.parse_args()

    # make directory to save outputs
    args.input_dir = args.input_dir.rstrip("/")  # remove trailing /
    dirname = os.path.dirname(args.input_dir)
    basename = os.path.basename(args.input_dir)
    output_dir = os.path.join(dirname, basename + f"+postprocess-v{args.version}")
    os.makedirs(output_dir, exist_ok=True)
    print(f"Saving outputs in {output_dir}...")

    # find all files
    # find all audio files
    filepaths = []
    for ext in ["flac", "wav", "mp3", "ogg"]:
        filepaths += glob.glob(os.path.join(args.input_dir, f"*.{ext}"))

    timings = []

    filepaths = filepaths[:20]

    # run algo on each audio file
    args = [(filepath, output_dir, args.version) for filepath in filepaths]
    with mp.Pool(32) as pool:
        pool.starmap(run, args)
