from tqdm import tqdm
import numpy as np
import re
from suno_utils.utils.text import read_jsonl, write_jsonl
import glob

TASK = "self_sim"

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/cover_filter/covers_raw.jsonl"
OUT_FILE = f"/app/suno/sara/ditto_v2_{TASK}_cover_scores.jsonl"

DITTO_SCORES


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_embeds(ditto_path):
    ditto_embeds = {}
    for file_path in tqdm(glob.glob(ditto_path)):
        ditto_data = np.load(file_path)
        for key in ditto_data:
            ditto_embeds[key] = np.mean(ditto_data[key], axis=0)

    print(f"Found {len(ditto_embeds)}")
    return ditto_embeds


print("Loading covers...")
cover_raw = read_jsonl(COVER_PATH, progress=True)


unique_fields = set()
for meta in tqdm(cover_raw):
    keys = list(meta.keys())
    for key in keys:
        unique_fields.add(key)
# print("Loading ditto locs...")
# ditto_embeds = get_ditto_embeds(DITTO_PATH)

filtered_cover_data = []
for meta in tqdm(cover_raw):
    source_title = meta["title"].lower()
    source_id = meta["id"]

    covers = meta["covers"]
    filtered_covers = []

    # parent_ditto = ditto_embeds[source_id]
    # parent_mean = np.mean(parent_ditto, axis=0)
    for cover in covers:
        cover_title = cover["title"].lower()
        cover_id = cover["id"]
        name_seems_match = (
            len(source_title) > 10
            and len(cover_title) > 10
            and len(set(source_title.lower()) & set(cover_title.lower())) >= 8
            and re.search(r"cover|remix", cover_title.lower())
            and not re.search(r"cover|remix", source_title.lower())
        )

        if name_seems_match and meta["views"] > 5000:
            filtered_covers.append(cover)
            """
            if cover_id not in ditto_embeds:
                print(f"{cover_id} not found")
                continue
            child_ditto = ditto_embeds[cover_id]
            child_mean = np.mean(child_ditto, axis=0)
            similarity = cosine_similarity(parent_mean, child_mean)

            if similarity < 0.9 and similarity > 0.3:
                filtered_cover = cover.copy()
                filtered_cover["parent_id"] = source_id
                filtered_cover["similarity"] = str(similarity)
                filtered_covers.append(filtered_cover)"
            """

        # if len(filtered_covers) >= 50:
        #    break

    if len(filtered_covers) > 0:
        # meta["centroid"] = [str(centroid) for centroid in source_centroid]
        source_data = {
            "id": meta["id"],
            "duration_s": meta["duration_s"],
            "s3_filepath": meta["s3_filepath"],
            "tags": meta["keywords"],
        }
        filtered_cover_data.append(source_data)
        for filtered in filtered_covers:
            cover_data = {
                "id": filtered["id"],
                "duration_s": filtered["duration_s"],
                "s3_filepath": filtered["s3_filepath"],
                "parent_id": meta["id"],
            }
            filtered_cover_data.append(cover_data)

write_jsonl(filtered_cover_data, "/app/suno/sara/filtered_cover_by_name_view.jsonl")
