import os
import json
import librosa
import soundfile as sf
import numpy as np
import glob
from pathlib import Path
from tqdm import tqdm
from werkzeug.utils import safe_join
from multiprocessing import Pool
import pandas as pd
from flask import Flask, render_template, jsonify, request, send_file

app = Flask(__name__)
base_dir = "/app/suno/christian/data/dpo_diffusion_test_set_1k/"
pair_dirs = glob.glob(os.path.join(base_dir, "*"))


def pair_similarity(pair, n_mfcc=20, hop_length=1024, n_fft=2048):
    audio_a, a_sr = sf.read(pair[0]["system_filepath"])
    audio_b, b_sr = sf.read(pair[1]["system_filepath"])

    # only compute with the first 60s
    audio_a = audio_a[: a_sr * 60]
    audio_b = audio_b[: b_sr * 60]

    # Compute MFCCs with custom parameters
    mfccs_a = librosa.feature.mfcc(
        y=audio_a.T, sr=a_sr, n_mfcc=n_mfcc, hop_length=hop_length, n_fft=n_fft
    )
    mfccs_b = librosa.feature.mfcc(
        y=audio_b.T, sr=b_sr, n_mfcc=n_mfcc, hop_length=hop_length, n_fft=n_fft
    )

    # compute mse between the two mfccs
    mse = np.mean((mfccs_a - mfccs_b) ** 2)

    # Normalize similarity to [0, 1] range
    # (cosine similarity normally ranges from -1 to 1)
    # similarity = np.exp(-0.1 * mse)

    return float(mse)


def get_audio_pairs():
    print("Getting audio pairs...")
    pairs = []
    for pair_dir in tqdm(pair_dirs):
        request_id = os.path.basename(pair_dir)
        audio_files = glob.glob(os.path.join(pair_dir, "*.mp3"))

        # Group files for this request_id
        pair = []
        for audio_file in audio_files:
            filename = os.path.basename(audio_file)
            item = {
                "system_filepath": audio_file,
                "filepath": f"/audio/{request_id}/{filename}",
                "s3_id": filename.split(".")[0],
                "request_id": request_id,
            }
            pair.append(item)

        if len(pair) == 2:  # Only add if we found files
            pairs.append(pair)

    # measure similarity of pairs in paralleized way
    with Pool(processes=64) as pool:
        similarities = pool.map(pair_similarity, pairs)

    # create a csv file with the filepaths and the similarities
    with open("similarities.csv", "w") as f:
        for pair, similarity in zip(pairs, similarities):
            f.write(
                f"{pair[0]['request_id']},{pair[0]['s3_id']},{pair[0]['filepath']},{pair[1]['s3_id']},{pair[1]['filepath']},{similarity}\n"
            )

    return pairs


# Initialize global variables
# check if csv exists
if os.path.exists("similarities.csv"):
    df = pd.read_csv(
        "similarities.csv",
        names=[
            "request_id",
            "s3_id_a",
            "filepath_a",
            "s3_id_b",
            "filepath_b",
            "similarity",
        ],
    )
    df = df.sort_values(by="similarity", ascending=False)
    # covert df to list of pairs
    pairs = df.to_dict(orient="records")
    print(f"Loaded {len(pairs)} pairs from csv")
else:
    pairs = get_audio_pairs()
    print(f"Loaded {len(pairs)} pairs")


@app.route("/")
def index():
    return render_template("index.html")


@app.route("/audio/<request_id>/<filename>")
def serve_audio(request_id, filename):
    try:
        full_path = safe_join(base_dir, request_id, filename)
        return send_file(full_path, mimetype="audio/mpeg")
    except Exception as e:
        return str(e), 404


@app.route("/get_pair/<int:pair_index>")
def get_pair(pair_index):
    if pair_index >= len(pairs):
        return jsonify({"status": "complete"})

    current_pair = pairs[pair_index]
    return jsonify(
        {
            "audio_a": {
                "id": current_pair["s3_id_a"],
                "url": current_pair["filepath_a"],
                "request_id": current_pair["request_id"],
            },
            "audio_b": {
                "id": current_pair["s3_id_b"],
                "url": current_pair["filepath_b"],
                "request_id": current_pair["request_id"],
            },
        }
    )


@app.route("/submit_response", methods=["POST"])
def submit_response():
    data = request.json
    result = {
        "pair_index": data["pairIndex"],
        "selected": data["selected"],
        "audio_a_id": data["audioAId"],
        "audio_b_id": data["audioBId"],
        "request_id": data["requestId"],
    }

    # Save results to file, append to existing file
    with open("results.jsonl", "a") as f:
        json.dump(result, f)
        f.write("\n")

    # print(f"Saved {len(results)} results to results.json")
    return jsonify({"status": "success"})


if __name__ == "__main__":
    app.run(debug=True)
