import json
import boto3
import tempfile
import os
import numpy as np
import statistics
from tqdm import tqdm
import sys
import torch
from suno_utils.audio import Audio
import torchaudio

sys.path.append("/home/sara/neon/ditto-training")
from ditto_v2.models.ditto import Ditto

S = torch.load("/app/suno/minz/models/ditto_cover.ckpt")
SS = {k.replace("model.", ""): v for k, v in S["state_dict"].items()}
ditto = Ditto(latent_dim=128, is_flash=False)
ditto.load_state_dict(SS)
ditto = ditto.eval()
ditto = ditto.cuda()
ditto.music_encoder = ditto.music_encoder.cuda()

TIMESTAMP = "2025_02_12-18_22_00"
DATA_FOLDER = "modal_runs"
OUTPUT_FOLDER = "scores"
DITTO_SR = 24000


def cosine_similarity(a, b):
    dot_product = np.dot(a, b)

    magnitude_a = np.sqrt(np.dot(a, a))
    magnitude_b = np.sqrt(np.dot(b, b))

    return dot_product / (magnitude_a * magnitude_b)


def get_audio_emb(audio):
    num_chunks = 4
    chunk_length_samples = 15 * 24000

    available_samples = max(0, len(audio) - chunk_length_samples)

    start_positions = (
        [int(i * available_samples / (num_chunks - 1)) for i in range(num_chunks)]
        if num_chunks > 1 and available_samples > 0
        else [0] * num_chunks
    )

    # Extract chunks
    chunks = []
    for start_pos in start_positions:
        chunk = audio[start_pos : start_pos + chunk_length_samples]
        if len(chunk) < chunk_length_samples:
            chunk = np.pad(chunk, (0, chunk_length_samples - len(chunk)))
        chunks.append(chunk)

    # Stack chunks and get embeddings
    chunks_tensor = torch.tensor(np.stack(chunks), device="cuda")

    emb = ditto.music_to_latent(chunks_tensor, "cover_sim").detach().cpu().numpy()

    return emb.mean(axis=0)


def get_task_scores(mappings, task):
    s3 = boto3.client("s3")
    bucket = "suno-data-uploads"
    folder_source = "studio/uploads/"
    folder_cover = folder_source  # f"tasks/feature_eval/cover_persona/{TIMESTAMP}/"

    scores = {}
    model = None
    for source_id, covers in tqdm(mappings.items()):
        scores[source_id] = []
        with tempfile.NamedTemporaryFile(suffix=".mp3") as temp_file:
            s3.download_file(
                bucket, os.path.join(folder_source, f"{source_id}.mp3"), temp_file.name
            )
            waveform, sr = torchaudio.load(temp_file.name)
            waveform = torch.mean(waveform, dim=0)
            if sr != DITTO_SR:
                resampler = torchaudio.transforms.Resample(sr, DITTO_SR)
            waveform = resampler(waveform)
            source_ditto = get_audio_emb(waveform)
        for cover_id in covers:
            s3.download_file(
                bucket, os.path.join(folder_cover, f"{cover_id}.mp3"), temp_file.name
            )
            waveform, sr = torchaudio.load(temp_file.name)
            waveform = torch.mean(waveform, dim=0)
            if sr != DITTO_SR:
                resampler = torchaudio.transforms.Resample(sr, DITTO_SR)
            waveform = resampler(waveform)
            child_ditto = get_audio_emb(waveform)

            score = cosine_similarity(source_ditto, child_ditto)
            scores[source_id].append((cover_id, score))

    return scores, "cover_sim", "filtered_45_hard"


def summarize_scores(scores):
    for key in scores:
        avg_score = statistics.mean(scores[key])
        min_score = min(scores[key])
        max_score = max(scores[key])
        print(f"{key}: average: {avg_score} min: {min_score} max: {max_score}")


if __name__ == "__main__":
    if not os.path.exists(OUTPUT_FOLDER):
        os.makedirs(OUTPUT_FOLDER)

    cover_path = os.path.join(DATA_FOLDER, f"cover_mappings_{TIMESTAMP}.json")
    score_cover = os.path.exists(cover_path)

    if score_cover:
        with open(cover_path) as f:
            cover_mappings = json.load(f)

        cover_scores, c_task, c_model = get_task_scores(cover_mappings, "cover_sim")

        print("COVER SCORES")
        # summarize_scores(cover_scores)
        cover_save_path = os.path.join(
            OUTPUT_FOLDER, f"feat_eval_cover_LABELED_{c_model}_{c_task}_{TIMESTAMP}.npz"
        )
        np.savez(cover_save_path, **cover_scores)
        print(f"Wrote cover scores to {cover_save_path}")
        print("\n")
