# ========= Subclustering + central representative + GPT summaries =========
import numpy as np
import pandas as pd
from sklearn.cluster import AgglomerativeClustering
from sklearn.metrics.pairwise import cosine_similarity
import sys

input_file = sys.argv[1]

# ---- SETTINGS ----
MIN_SIZE_FOR_SUB = 20            # prag: doar clusterele cu >= acest număr vor fi subclusterizate
TOP_N_EXAMPLES_FOR_SUMMARY = 5   # câte exemple să trimiți la GPT pt. sumar
USE_GPT_SUMMARY = True           # setează False dacă nu vrei sumar GPT
MAIN_DISTANCE_THRESHOLD = 0.30   # dacă vrei să refaci clusteringul principal mai lax/strict
SUB_DISTANCE_THRESHOLD = 0.20    # subclusterizare mai strictă (sim cos > ~0.80)
OUT_FILE = f'{input_file.split(".csv")[0]}_painpoint_clusters_enhanced.csv'

# ---- 1) construim un dicționar canonical -> embedding (din runda ta de embedding)
#      PRESUPUNE că 'texts' sunt în aceeași ordine ca 'embeddings'
canon2emb = {t: embeddings[i] for i, t in enumerate(texts)}

# ---- 2) funcții utilitare ----
def central_representative(canon_list):
    """returnează (repr_text, centrality_score, centroid_vector) pt. un grup."""
    vecs = np.vstack([canon2emb[c] for c in canon_list])
    centroid = vecs.mean(axis=0, keepdims=True)
    sims = cosine_similarity(vecs, centroid).ravel()
    idx = int(np.argmax(sims))
    return canon_list[idx], float(sims[idx]), centroid.ravel()

def agglom_subcluster(canon_list, distance_threshold):
    """subclusterizează o listă de canonic-uri cu un prag pe distanța cosinus"""
    vecs = np.vstack([canon2emb[c] for c in canon_list])
    # direct pe embeddings, metric="cosine"
    sub = AgglomerativeClustering(
        n_clusters=None,
        metric="cosine",
        linkage="average",
        distance_threshold=distance_threshold
    ).fit(vecs)
    return sub.labels_

def gpt_summary(title, examples):
    """scurt sumar GPT al pain point-ului din câteva exemple"""
    from openai import OpenAI
    client = OpenAI(api_key=os.getenv("OPENAI_API_KEY"))
    prompt = (
        "You will summarize a set of short user pain point statements into a single, "
        "concise problem statement (max 1 sentence) and 3 bullet insights. "
        "Focus on what's hard for users and why. Avoid solutions.\n\n"
        f"Title: {title}\n"
        "Examples:\n- " + "\n- ".join(examples[:TOP_N_EXAMPLES_FOR_SUMMARY])
    )
    resp = client.chat.completions.create(
        model="gpt-4o-mini",
        temperature=0.2,
        messages=[
            {"role": "system", "content": "You are a precise UX researcher."},
            {"role": "user", "content": prompt}
        ]
    )
    return resp.choices[0].message.content.strip()

# ---- 3) pregătim un tabel de lucru din clusteringul tău curent ----
df_clusters = df_clusters.copy()
sizes = df_clusters["cluster"].value_counts().to_dict()
df_clusters["cluster_size"] = df_clusters["cluster"].map(sizes)

# reprezentativ central la nivel de cluster
cluster_repr = {}
cluster_centroid = {}
cluster_repr_centrality = {}
for cid, grp in df_clusters.groupby("cluster"):
    repr_text, centr, centroid_vec = central_representative(grp["canonical"].tolist())
    cluster_repr[cid] = repr_text
    cluster_centroid[cid] = centroid_vec
    cluster_repr_centrality[cid] = centr

df_clusters["cluster_representative"] = df_clusters["cluster"].map(cluster_repr)
df_clusters["cluster_repr_centrality"] = df_clusters["cluster"].map(cluster_repr_centrality)

# ---- 4) subclusterizare doar pe clusterele mari ----
sub_rows = []
for cid, grp in df_clusters.groupby("cluster", sort=False):
    canons = grp["canonical"].tolist()
    if len(canons) >= MIN_SIZE_FOR_SUB:
        sub_labels = agglom_subcluster(canons, SUB_DISTANCE_THRESHOLD)
    else:
        sub_labels = np.zeros(len(canons), dtype=int)  # un singur subcluster „0”

    # adaugă înregistrări cu subcluster id
    for canon, sublab in zip(canons, sub_labels):
        sub_rows.append((cid, canon, int(sublab)))

df_sub = pd.DataFrame(sub_rows, columns=["cluster", "canonical", "subcluster"])

# merge în tabelul principal
df_enh = df_clusters.merge(df_sub, on=["cluster", "canonical"], how="left")

# dimensiuni subcluster
sub_sizes = df_enh.groupby(["cluster", "subcluster"]).size().rename("subcluster_size").reset_index()
df_enh = df_enh.merge(sub_sizes, on=["cluster", "subcluster"], how="left")

# reprezentativ central pe subcluster
sub_repr_map = {}
sub_centrality_map = {}
for (cid, scid), grp in df_enh.groupby(["cluster", "subcluster"]):
    repr_text, centr, _ = central_representative(grp["canonical"].tolist())
    sub_repr_map[(cid, scid)] = repr_text
    sub_centrality_map[(cid, scid)] = centr

df_enh["subcluster_representative"] = df_enh.apply(
    lambda r: sub_repr_map[(r["cluster"], r["subcluster"])], axis=1
)
df_enh["sub_repr_centrality"] = df_enh.apply(
    lambda r: sub_centrality_map[(r["cluster"], r["subcluster"])], axis=1
)

# ---- 5) sumar GPT (opțional) pentru fiecare (sub)cluster ----
cluster_summary = {}
subcluster_summary = {}
if USE_GPT_SUMMARY:
    # sumar pe cluster
    for cid, grp in df_enh.groupby("cluster"):
        exs = grp["canonical"].sample(min(TOP_N_EXAMPLES_FOR_SUMMARY, len(grp)), random_state=42).tolist()
        title = cluster_repr[cid]
        cluster_summary[cid] = gpt_summary(title, exs)

    # sumar pe subcluster
    for (cid, scid), grp in df_enh.groupby(["cluster", "subcluster"]):
        exs = grp["canonical"].sample(min(TOP_N_EXAMPLES_FOR_SUMMARY, len(grp)), random_state=42).tolist()
        title = sub_repr_map[(cid, scid)]
        subcluster_summary[(cid, scid)] = gpt_summary(title, exs)

    df_enh["cluster_summary"] = df_enh["cluster"].map(cluster_summary)
    df_enh["subcluster_summary"] = df_enh.apply(
        lambda r: subcluster_summary[(r["cluster"], r["subcluster"])], axis=1
    )

# ---- 6) sortare: mai întâi clusterele mari, apoi subclusterele mari ----
df_enh = df_enh.sort_values(
    ["cluster_size", "subcluster_size", "cluster", "subcluster", "canonical"],
    ascending=[False, False, True, True, True]
)

# ---- 7) salvare ----
cols = [
    "cluster", "cluster_size", "cluster_representative", "cluster_repr_centrality",
    "subcluster", "subcluster_size", "subcluster_representative", "sub_repr_centrality",
    "canonical"
]
if USE_GPT_SUMMARY:
    cols = cols + ["cluster_summary", "subcluster_summary"]

df_enh[cols].to_csv(OUT_FILE, index=False)
print(f"\nSaved enhanced results → {OUT_FILE}")

# ---- 8) print scurt pentru verificare ----
print("\n=== Top clusters (with subclusters) ===")
for cid, grp in df_enh.groupby("cluster", sort=False):
    csize = int(grp["cluster_size"].iloc[0])
    crepr = grp["cluster_representative"].iloc[0]
    print(f"\nCluster {cid} ({csize} items) | Repr: {crepr}")
    for (cid2, scid), g2 in grp.groupby(["cluster", "subcluster"], sort=False):
        ssz = int(g2["subcluster_size"].iloc[0])
        srepr = g2["subcluster_representative"].iloc[0]
        print(f"  - Sub {scid} ({ssz}) | Repr: {srepr}")
        for ex in g2["canonical"].head(3).tolist():
            print(f"      · {ex}")
