from flask import Flask, render_template, request, redirect, url_for
import os
import json
import re
import numpy as np
from tqdm import tqdm

app = Flask(__name__)

# Path to where audio files are stored
DATA_NAME = 'v3-distill-data-t0'
AUDIO_BASE_DIR = f'static/{DATA_NAME}'
RESULTS_FILE = 'results.json'

# Load or initialize results
if os.path.exists(RESULTS_FILE):
    with open(RESULTS_FILE, 'r') as f:
        results = json.load(f)
else:
    results = {}

def get_all_pairs():
    pairs = []
    for dir_name in tqdm(sorted(os.listdir(AUDIO_BASE_DIR))):
        dir_path = os.path.join(AUDIO_BASE_DIR, dir_name)
        if not os.path.isdir(dir_path):
            continue

        files = os.listdir(dir_path)
        a_files = [f for f in files if re.search(r'_a\.mp3$', f)]
        b_files = [f for f in files if re.search(r'_b\.mp3$', f)]

        for a_file in a_files:
            base = a_file.replace('_a.mp3', '')
            b_file = f"{base}_b.mp3"
            if b_file in b_files:
                model_name = base.split(f"{dir_name}_")[1]  # get model suffix
                try:
                    a_metadata = np.load(os.path.join(dir_path, f"{dir_name}_{model_name}_a__metadata.npz"), allow_pickle=True)
                    b_metadata = np.load(os.path.join(dir_path, f"{dir_name}_{model_name}_b__metadata.npz"), allow_pickle=True)

                    a_meta = {k: a_metadata[k].tolist() for k in a_metadata}
                    b_meta = {k: b_metadata[k].tolist() for k in b_metadata}
                except Exception as e:
                    print(f"Metadata load failed for {dir_name}: {e}")
                    a_meta, b_meta = {}, {}

                pairs.append({
                    'pair_id': dir_name,
                    'a_path': os.path.join(DATA_NAME, dir_name, a_file),
                    'b_path': os.path.join(DATA_NAME, dir_name, b_file),
                    'a_meta': a_meta,
                    'b_meta': b_meta,
                })
    return pairs


# Load all available pairs
all_pairs = get_all_pairs()




@app.route('/')
def index():
    """Redirect to the next unlabeled pair, or show 'done'."""
    for pair in all_pairs:
        pair_id = pair['pair_id']
        if pair_id not in results:
            return redirect(url_for('label_pair', pair_id=pair_id))
    return "✅ All audio pairs labeled!"

@app.route('/label/<pair_id>')
def label_pair(pair_id):
    """Display A/B audio and label options."""
    labeled_count = len(results)
    total_count = len(all_pairs)
    pair = next((p for p in all_pairs if p['pair_id'] == pair_id), None)
    if not pair:
        return "❌ Pair not found", 404
    return render_template(
        'label.html',
        pair_id=pair_id,
        audio_a=pair['a_path'],
        audio_b=pair['b_path'],
        a_meta=pair['a_meta'],
        b_meta=pair['b_meta'],
        labeled_count=labeled_count,
        total_count=total_count
    )
@app.route('/submit', methods=['POST'])
def submit():
    """Save choice for a pair and redirect."""
    pair_id = request.form['pair_id']
    choice = request.form['choice']  # 'A' or 'B'

    results[pair_id] = choice

    with open(RESULTS_FILE, 'w') as f:
        json.dump(results, f, indent=2)

    return redirect(url_for('index'))

if __name__ == '__main__':
    app.run(debug=True, port=8080)
