import json
import numpy as np
from tqdm import tqdm
import glob

TASK = "self_sim"  # update this


DATA_PATH = f"/app/suno/sara/ditto_v2_{TASK}/metas/*.jsonl"
DITTO_PATH = f"/app/suno/sara/ditto_v2_{TASK}/*.npz"
COVER_PATH = "/app/suno/sara/metas_v0_cover.jsonl"
OUT_FILE = f"/app/suno/sara/ditto_v2_{TASK}_cover_scores.jsonl"


def load_jsonl_cover(data_path):
    covers = {}  # parent id -> [child ids]

    PARENT = "parent_id"
    ID = "id"

    with open(data_path, "r", encoding="utf-8") as file:
        for line in tqdm(file):
            try:
                sample = json.loads(line)
                if PARENT in sample:
                    parent = sample[PARENT]
                    if parent not in covers:
                        covers[parent] = []
                    covers[parent].append(sample[ID])

            except json.JSONDecodeError as e:
                print(f"Error decoding JSON: {e}")

    return covers


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_ditto_file_locs(ditto_path):
    ditto_locs = {}
    for file_path in tqdm(glob.glob(ditto_path)):
        ditto_data = np.load(file_path)
        for key in ditto_data:
            ditto_locs[key] = file_path

    return ditto_locs


if __name__ == "__main__":
    print("Parsing cover data...")
    covers = load_jsonl_cover(COVER_PATH)

    print("Parsing ditto data...")
    ditto_locs = get_ditto_file_locs(DITTO_PATH)

    ditto_scores = []
    misses = 0
    entries = 0

    print("Calculating scores...")
    for parent in tqdm(covers):
        child_tracks = covers[parent]
        if parent not in ditto_locs:
            misses += 1
            continue
        parent_ditto = np.load(ditto_locs[parent])[parent]
        parent_mean = np.mean(parent_ditto, axis=0)
        for child in child_tracks:
            if child not in ditto_locs:
                misses += 1
                continue
            child_ditto = np.load(ditto_locs[child])[child]
            child_mean = np.mean(child_ditto, axis=0)
            score = cosine_similarity(parent_mean, child_mean)
            ditto_scores.append({"parent_id": parent, "child_id": child, f"score_{TASK}": float(score)})
        entries += 1
        # if entries > 10000:
        #    break

    num_pairs = len(ditto_scores)
    print(f"Done processing {num_pairs} pairs")
    print(f"{misses} missed ids when processing ditto covers, thats {misses/num_pairs * 100}%")

    print(f"Writing results to {OUT_FILE}...")
    try:
        with open(OUT_FILE, "w", encoding="utf-8") as f:
            for item in ditto_scores:
                json_line = json.dumps(item, ensure_ascii=False)
                f.write(json_line + "\n")
    except Exception as e:
        raise Exception(f"Error writing to JSONL file: {str(e)}")
    print("Done!")
