#!/usr/bin/env python3
"""
JTBD Tag Cloud Generator — Auto Theming with Embeddings
-------------------------------------------------------

Features:
- Lemmatization + custom stopwords
- Min frequency filter
- Relative frequency (%)
- Embeddings via sentence-transformers
- Clustering: KMeans (fixed k) or HDBSCAN (auto)
- Auto-theme labeling (top terms) + exemplar n-grams (closest to centroid)
- Saves:
    - <col>_top_<n>grams.csv            (raw n-grams with counts & %)
    - <col>_themes_<n>grams.csv         (themes summary)
    - <col>_themes_assignments_<n>.csv  (each n-gram → theme)

Usage examples:
  python jtbd_tag_cloud.py -i jtbd_openai_results.csv --ngrams 1 2 3 --min-count 2 \
    --cluster-method kmeans --num-themes 8

  python jtbd_tag_cloud.py -i jtbd_openai_results.csv --cluster-method hdbscan
"""

import argparse
import pandas as pd
from collections import Counter, defaultdict
from pathlib import Path
import re
import numpy as np
import unicodedata
import math

# Text normalization
import nltk
from nltk.stem import WordNetLemmatizer

# Embeddings & clustering
from sentence_transformers import SentenceTransformer
from sklearn.cluster import KMeans
from sklearn.feature_extraction.text import CountVectorizer
from sklearn.metrics.pairwise import cosine_distances

# HDBSCAN is optional; only used if selected
try:
    import hdbscan
    HAS_HDBSCAN = True
except Exception:
    HAS_HDBSCAN = False


def parse_args():
    p = argparse.ArgumentParser(
        description="Generate top n-grams from JTBD-labeled comments and auto-cluster into themes."
    )
    p.add_argument("-i", "--input", required=True, help="Path to input CSV file")
    p.add_argument("--ngrams", type=int, nargs="+", default=[1, 2, 3],
                   help="List of n-gram sizes to compute (e.g., 1 2 3)")
    p.add_argument("--topn", type=int, default=50, help="Max n-grams per n-size to keep (after filtering)")
    p.add_argument("--min-count", type=int, default=2, help="Minimum count to keep an n-gram")
    p.add_argument("--stopwords", type=str, default=None, help="Optional stopwords file (one per line)")
    p.add_argument("--embedding-model", type=str, default="sentence-transformers/all-MiniLM-L6-v2",
                   help="Sentence-Transformers model name")
    p.add_argument("--cluster-method", choices=["kmeans", "hdbscan"], default="kmeans",
                   help="Clustering algorithm")
    p.add_argument("--num-themes", type=int, default=8,
                   help="Number of themes (KMeans only)")
    p.add_argument("--columns", nargs="+", default=["situation", "struggle", "outcome"],
                   help="Columns to analyze")
    return p.parse_args()


def ensure_nltk():
    # Make lemmatizer data available without crashing if missing
    try:
        nltk.data.find("corpora/wordnet")
    except LookupError:
        nltk.download("wordnet", quiet=True)
    try:
        nltk.data.find("corpora/omw-1.4")
    except LookupError:
        nltk.download("omw-1.4", quiet=True)


def load_stopwords(extra_file=None):
    base = set("""
    the and to a of in for on at with an is it that this so be as by or are was were
    from but have has had when i me my you your they their them we us our just really
    like get got make makes made much many very more less lot lots kind sort maybe
    """.split())
    if extra_file and Path(extra_file).exists():
        with open(extra_file, "r", encoding="utf-8") as f:
            base |= {line.strip().lower() for line in f if line.strip()}
    return base


def clean_text(text: str) -> str:
    """Lowercase, normalize accents, keep letters and spaces."""
    if not isinstance(text, str):
        return ""
    text = text.lower()
    text = unicodedata.normalize("NFKD", text)
    text = re.sub(r"[^a-z\s]", " ", text)
    return re.sub(r"\s+", " ", text).strip()


lemmatizer = WordNetLemmatizer()


def tokenize(text: str, stopwords):
    tokens = []
    for t in text.split():
        if t in stopwords or len(t) < 3:
            continue
        lemma = lemmatizer.lemmatize(t)
        if lemma not in stopwords and len(lemma) >= 3:
            tokens.append(lemma)
    return tokens


def build_ngrams(tokens_list, n=1):
    ngrams = []
    for tokens in tokens_list:
        for i in range(len(tokens) - n + 1):
            ngrams.append(" ".join(tokens[i:i + n]))
    return ngrams


def top_ngrams(tokens_list, n=1, topn=50, min_count=2):
    ngrams = build_ngrams(tokens_list, n)
    counter = Counter(ngrams)
    items = [(gram, cnt) for gram, cnt in counter.items() if cnt >= min_count]
    items.sort(key=lambda x: x[1], reverse=True)

    # limit to topn *after* filtering
    items = items[:topn] if topn else items

    total = sum(cnt for _, cnt in items) or 1
    results = [(gram, cnt, round(cnt / total * 100, 2)) for gram, cnt in items]
    return results


def embed_texts(texts, model_name):
    model = SentenceTransformer(model_name)
    emb = model.encode(texts, batch_size=64, show_progress_bar=False, normalize_embeddings=True)
    return emb


def cluster_kmeans(embeddings, k):
    km = KMeans(n_clusters=k, n_init="auto", random_state=42)
    labels = km.fit_predict(embeddings)
    centers = km.cluster_centers_
    return labels, centers


def cluster_hdbscan(embeddings):
    if not HAS_HDBSCAN:
        raise RuntimeError("HDBSCAN not installed. Please `pip install hdbscan` or use --cluster-method kmeans.")
    clusterer = hdbscan.HDBSCAN(min_cluster_size=5, metric="euclidean")  # embeddings are normalized
    labels = clusterer.fit_predict(embeddings)
    # Compute pseudo-centers as mean of members for labeling/exemplars
    centers = []
    for c in sorted(set([l for l in labels if l != -1])):
        idx = (labels == c)
        centers.append(embeddings[idx].mean(axis=0))
    # Map cluster id → center index
    cluster_to_center_idx = {}
    for i, c in enumerate(sorted(set([l for l in labels if l != -1]))):
        cluster_to_center_idx[c] = i
    return labels, centers, cluster_to_center_idx


def auto_label_themes(ngrams, labels, stopwords, kmeans_centers=None, embeddings=None, hdbscan_center_map=None):
    """
    ngrams: list[str]
    labels: array[int] (cluster label per ngram; -1 may indicate noise)
    Returns:
      theme_summary: list of dicts (theme, count, percent, label_terms, exemplars)
      assignments_df: DataFrame with ngram, count, percent, theme_id, theme_label
    """
    data = []
    for (gram, cnt, pct), lab in zip(ngrams, labels):
        data.append({"ngram": gram, "count": cnt, "percent": pct, "cluster": lab})
    df = pd.DataFrame(data)

    # Separate noise (cluster = -1) if present
    clusters = sorted([c for c in df["cluster"].unique() if c != -1])
    if not clusters:
        # All noise or empty — bail gracefully
        return [], pd.DataFrame(data)

    # Build a vocabulary over n-grams for labeling (bag of words across n-grams per cluster)
    # We’ll pick top frequent terms per cluster as the label
    vectorizer = CountVectorizer(stop_words=list(stopwords))
    # Represent n-gram texts as bag-of-words (not their counts in corpus)
    bow = vectorizer.fit_transform(df["ngram"])
    vocab = {i: t for t, i in vectorizer.vocabulary_.items()}
    inv_vocab = {i: t for i, t in vocab.items()}

    theme_rows = []
    label_map = {}

    # Precompute distances if we have centers (for exemplars)
    exemplar_by_cluster = defaultdict(list)
    if embeddings is not None:
        # For KMeans
        if kmeans_centers is not None:
            dists = cosine_distances(embeddings, kmeans_centers)
            # dists[i, c] is distance of n-gram i to center c
            for idx, c in enumerate(clusters):
                # exemplar = n-grams with smallest distance to center c
                cluster_indices = df.index[df["cluster"] == c].tolist()
                if not cluster_indices:
                    continue
                # Map to distances
                ranked = sorted(cluster_indices, key=lambda i: dists[i, c])
                exemplar_by_cluster[c] = ranked[:5]
        # For HDBSCAN pseudo-centers
        elif hdbscan_center_map is not None:
            # Build center matrix in cluster id order
            centers = []
            c_order = []
            for c in clusters:
                centers.append(hdbscan_center_map[c])
                c_order.append(c)
            centers = np.vstack(centers)
            dists = cosine_distances(embeddings, centers)
            for idx, c in enumerate(clusters):
                cluster_indices = df.index[df["cluster"] == c].tolist()
                ranked = sorted(cluster_indices, key=lambda i: dists[i, idx])
                exemplar_by_cluster[c] = ranked[:5]

    for c in clusters:
        mask = (df["cluster"] == c)
        dfc = df[mask]
        total_cnt = int(dfc["count"].sum())
        total_cnt_all = int(df[df["cluster"] != -1]["count"].sum() or 1)
        percent = round(total_cnt / total_cnt_all * 100, 2)

        # Top terms for auto-label
        # Sum BoW rows for this cluster
        row_indices = dfc.index.tolist()
        if row_indices:
            cluster_bow = bow[row_indices].sum(axis=0).A1  # 1D array
            # get top 3-5 terms
            top_idx = cluster_bow.argsort()[::-1][:5]
            label_terms = [inv_vocab[i] for i in top_idx if cluster_bow[i] > 0]
        else:
            label_terms = []

        # Human-friendly label from top terms
        if label_terms:
            label = ", ".join(label_terms[:3])
        else:
            label = f"Theme {c}"

        label_map[c] = label

        # Exemplars (closest to center) or just top by count
        if exemplar_by_cluster.get(c):
            ex_idx = exemplar_by_cluster[c]
            exemplars = df.loc[ex_idx, "ngram"].tolist()
        else:
            exemplars = dfc.sort_values("count", ascending=False)["ngram"].head(5).tolist()

        theme_rows.append({
            "theme_id": c,
            "theme_label": label,
            "count": total_cnt,
            "percent": percent,
            "top_terms": label_terms,
            "exemplars": exemplars
        })

    theme_rows.sort(key=lambda r: r["count"], reverse=True)

    # Add theme label to assignments
    df["theme_id"] = df["cluster"]
    df["theme_label"] = df["theme_id"].map(label_map).fillna("Noise")

    return theme_rows, df[["ngram", "count", "percent", "theme_id", "theme_label"]]


def analyze_column(col_name, df_yes, args, stopwords):
    texts = df_yes[col_name].dropna().tolist()
    cleaned = [clean_text(t) for t in texts]
    tokenized = [tokenize(t, stopwords) for t in cleaned]

    print(f"\n================ {col_name.upper()} ================")
    all_theme_summaries = []

    for n in args.ngrams:
        # 1) n-grams + counts
        results = top_ngrams(tokenized, n=n, topn=args.topn, min_count=args.min_count)
        out_dir = Path(args.input).parent

        # Save raw n-grams
        df_out = pd.DataFrame(results, columns=[f"{n}-gram", "count", "percent"])
        out_file = out_dir / f"{col_name}_top_{n}grams.csv"
        df_out.to_csv(out_file, index=False)
        print(f"Saved n-grams → {out_file}")

        if len(results) == 0:
            print(f"(No n-grams kept for {n}-grams at min-count={args.min_count})")
            continue

        # 2) Embeddings for n-gram strings
        ngram_texts = [row[0] for row in results]
        try:
            emb = embed_texts(ngram_texts, args.embedding_model)
        except Exception as e:
            print(f"Embedding failed ({e}). Skipping theming for {n}-grams.")
            continue

        # 3) Clustering
        labels = None
        kmeans_centers = None
        hdbscan_center_map = None

        if args.cluster_method == "kmeans":
            k = max(2, args.num_themes)  # ensure >=2
            labels, kmeans_centers = cluster_kmeans(emb, k)
        else:
            if not HAS_HDBSCAN:
                print("HDBSCAN not installed; falling back to kmeans.")
                labels, kmeans_centers = cluster_kmeans(emb, max(2, args.num_themes))
            else:
                labels, centers, c2idx = cluster_hdbscan(emb)
                # build map cluster id -> center vector for exemplar selection
                hdbscan_center_map = {}
                i_order = 0
                for c in sorted(set([l for l in labels if l != -1])):
                    hdbscan_center_map[c] = centers[i_order]
                    i_order += 1

        # 4) Auto-label themes & exemplars
        theme_rows, assignments = auto_label_themes(
            results,
            labels,
            stopwords,
            kmeans_centers=kmeans_centers,
            embeddings=emb,
            hdbscan_center_map=hdbscan_center_map
        )

        if not theme_rows:
            print("No clusters found (or all noise). Skipping theme tables.")
            continue

        # Save themes summary
        df_themes = pd.DataFrame(theme_rows)
        out_file_themes = out_dir / f"{col_name}_themes_{n}grams.csv"
        df_themes.to_csv(out_file_themes, index=False)
        print(f"Saved themes summary → {out_file_themes}")

        # Save assignments
        out_file_assign = out_dir / f"{col_name}_themes_assignments_{n}grams.csv"
        assignments.to_csv(out_file_assign, index=False)
        print(f"Saved theme assignments → {out_file_assign}")

        # Print concise preview
        preview = df_themes[["theme_label", "count", "percent"]].head(10)
        print("\nTop themes:")
        print(preview.to_string(index=False))

        all_theme_summaries.append((n, df_themes))

    return all_theme_summaries


def main():
    args = parse_args()
    ensure_nltk()
    stopwords = load_stopwords(args.stopwords)

    input_path = Path(args.input)
    df = pd.read_csv(input_path)

    # Basic sanity checks
    for col in args.columns:
        if col not in df.columns:
            raise ValueError(f"Input CSV missing required column: '{col}'")
    if "jtbd" not in df.columns or "comment" not in df.columns:
        raise ValueError("Input CSV must include 'jtbd' and 'comment' columns.")

    # Filter JTBD=YES and unique comments
    df_yes = df[df["jtbd"] == "YES"].drop_duplicates(subset=["comment"])

    # Analyze requested columns
    for col in args.columns:
        analyze_column(col, df_yes, args, stopwords)


if __name__ == "__main__":
    main()
