import pandas as pd
from keybert import KeyBERT
from tqdm import tqdm
from collections import Counter
import re

# ---------------- CONFIG ----------------
INPUT_FILE = "jtbd_semantic_matches.csv"   # fișierul tău cu rezultate din embeddings
OUTPUT_FILE = "jtbd_presence_queries.csv"
SIM_THRESHOLD = 0.55                  # scor minim pentru a considera comentariul relevant
TOP_KEYWORDS = 15                     # câte expresii să extragem per JTBD
# ----------------------------------------

print("🔹 Loading data...")
df = pd.read_csv(INPUT_FILE)

# filtrează doar comentariile relevante
df = df[df["similarity_score"] >= SIM_THRESHOLD]

print(f"🔹 Using {len(df)} relevant comment matches (score ≥ {SIM_THRESHOLD})")

# agregare pe JTBD
grouped = df.groupby("jtbd_id")

kw_model = KeyBERT('all-MiniLM-L6-v2')

results = []

for jtbd_id, group in tqdm(grouped, desc="Processing JTBDs"):
    jtbd_text = group["original_text"].iloc[0] if "original_text" in group else "N/A"
    comments = group["similar_comment"].dropna().tolist()

    # calculează scoruri de prezență
    mean_sim = group["similarity_score"].mean()
    max_sim = group["similarity_score"].max()
    count = len(group)
    presence_score = mean_sim * count

    # unește toate comentariile relevante
    text_blob = " ".join(comments)
    text_blob = re.sub(r"http\S+|www\S+|[\r\n]+", " ", text_blob)

    # extrage expresii semnificative
    try:
        keywords = kw_model.extract_keywords(
            text_blob,
            keyphrase_ngram_range=(1, 3),
            stop_words='english',
            top_n=TOP_KEYWORDS
        )
        phrases = [kw for kw, _ in keywords]
    except Exception:
        phrases = []

    # generează query pentru Reddit
    reddit_query = " OR ".join([f'"{p}"' for p in phrases])

    results.append({
        "jtbd_id": jtbd_id,
        "jtbd_text": jtbd_text,
        "count_relevant_comments": count,
        "mean_similarity": round(mean_sim, 3),
        "max_similarity": round(max_sim, 3),
        "presence_score": round(presence_score, 3),
        "key_phrases": ", ".join(phrases),
        "reddit_query": reddit_query
    })

# salvare rezultate
out_df = pd.DataFrame(results)
out_df.sort_values("presence_score", ascending=False, inplace=True)
out_df.to_csv(OUTPUT_FILE, index=False)

print(f"✅ Done! Saved {len(out_df)} JTBD presence results to {OUTPUT_FILE}")
