import pandas as pd
import re
import nltk
from sentence_transformers import SentenceTransformer
import hdbscan
import numpy as np
from sklearn.feature_extraction.text import TfidfVectorizer

nltk.download('punkt')

def split_sentences(comment):
    return nltk.sent_tokenize(comment)

def is_junk_sentence(sentence):
    """Drop single-sentence fragments shorter than 5 words."""
    return len(sentence.strip().split()) < 5

def passes_structural_filters(text):
    if is_junk_sentence(text):
        return False
    words = text.split()
    if len(words) < 10 or len(words) > 150:
        return False
    if not re.search(r"\b(I|me|my|myself)\b", text, re.IGNORECASE):
        return False
    markers = [
        r"\bbecause\b", r"\bsince\b", r"\bdue to\b", r"\bas a result\b",
        r"\bwhen\b", r"\bafter\b", r"\bbefore\b", r"\bwhile\b", r"\bwhenever\b",
        r"\bso I can\b", r"\bso that\b", r"\bin order to\b",
        r"\bhelps me\b", r"\ballows me\b", r"\benables me\b", r"\blets me\b"
    ]
    if not any(re.search(m, text, re.IGNORECASE) for m in markers):
        return False
    verb_like = re.findall(r"\b\w+(ed|ing)\b", text.lower())
    if len(verb_like) < 2:
        return False
    if text.strip().endswith("?"):
        return False
    return True

def keep_comment(comment, drivers):
    if is_junk_sentence(comment):
        return False

    sentences = split_sentences(comment)
    structural_valid = any(passes_structural_filters(s) for s in sentences)

    high_signal = {
        "realization","curiosity","confusion","annoyance",
        "disapproval","disappointment","sadness","fear","remorse","anxiety","neutral"
    }
    low_signal = {"gratitude","admiration","approval","love","amusement"}

    drivers = {d.strip().lower() for d in str(drivers).split(",") if d.strip() and d.lower() != "nan"}
    emotion_valid = bool(drivers & high_signal)
    low_only = drivers.issubset(low_signal) if drivers else False

    return (structural_valid or emotion_valid) and not low_only

def extract_valid_sentences(comment):
    return [s for s in split_sentences(comment) if passes_structural_filters(s)]

def cluster_comments(comments):
    print("Encoding comments...")
    model = SentenceTransformer("all-MiniLM-L6-v2")
    embeddings = model.encode(comments, show_progress_bar=True)

    print("Clustering...")
    clusterer = hdbscan.HDBSCAN(min_cluster_size=25, min_samples=5, metric="euclidean")
    labels = clusterer.fit_predict(embeddings)

    return labels

def get_top_keywords(cluster_texts, top_n=10):
    vec = TfidfVectorizer(stop_words="english", max_features=5000, ngram_range=(1,2))
    X = vec.fit_transform(cluster_texts)
    sums = np.array(X.sum(axis=0)).flatten()
    terms = vec.get_feature_names_out()
    top_idx = sums.argsort()[::-1][:top_n]
    return [terms[i] for i in top_idx]

import pandas as pd
import numpy as np
from sentence_transformers import SentenceTransformer
import hdbscan
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.metrics.pairwise import cosine_similarity
from collections import Counter

# ========= Helpers ========= #

def get_top_keywords(docs, n=10):
    """Extract top keywords for a cluster using TF-IDF."""
    if len(docs) == 0:
        return []
    vectorizer = TfidfVectorizer(stop_words="english", ngram_range=(1, 2), min_df=2)
    X = vectorizer.fit_transform(docs)
    scores = np.asarray(X.sum(axis=0)).ravel()
    keywords = [(word, scores[idx]) for word, idx in vectorizer.vocabulary_.items()]
    keywords = sorted(keywords, key=lambda x: x[1], reverse=True)[:n]
    return [word for word, score in keywords]


def embed_sentences(sentences, model_name="all-mpnet-base-v2"):
    model = SentenceTransformer(model_name)
    return model.encode(sentences, show_progress_bar=True)


def cluster_sentences(embeddings, min_cluster_size=10):
    clusterer = hdbscan.HDBSCAN(min_cluster_size=min_cluster_size,
                                metric="euclidean",
                                cluster_selection_method="eom")
    cluster_labels = clusterer.fit_predict(embeddings)
    return cluster_labels


# ========= Main ========= #

def main(input_file, output_file):
    # Load file
    df = pd.read_csv(input_file)
    df['comment'] = df['comment'].astype(str).str.strip()
    df['emotional_drivers'] = df['emotional_drivers'].astype(str).str.lower().str.strip()

    # Filter to JTBD-worthy rows (if you already have a keep_comment function)
    if "keep_comment" in globals():
        df = df[df.apply(lambda row: keep_comment(row["comment"], row["emotional_drivers"]), axis=1)]

    # Split into JTBD sentences
    if "extract_valid_sentences" in globals():
        df["jtbd_sentences"] = df["comment"].apply(extract_valid_sentences)
    else:
        df["jtbd_sentences"] = df["comment"].apply(lambda c: [c])  # fallback

    # Flatten into sentence records
    sentence_records = []
    for idx, row in df.iterrows():
        for sent in row["jtbd_sentences"]:
            sentence_records.append({
                "comment_id": idx,
                "comment": row["comment"],
                "sentence": sent
            })
    sentences_df = pd.DataFrame(sentence_records)
    print(f"Found {len(sentences_df)} JTBD sentences to cluster.")

    # No longer needed as we cluster in a different place
    # # Embeddings
    # embeddings = embed_sentences(sentences_df["sentence"].tolist())

    # # Clustering
    # topics = cluster_sentences(embeddings, min_cluster_size=10)
    # sentences_df["topic"] = topics

    # # Label clusters
    # cluster_keywords = {}
    # for cluster_id in set(topics):
    #     if cluster_id == -1:
    #         continue
    #     cluster_sents = sentences_df[sentences_df["topic"] == cluster_id]["sentence"].tolist()
    #     cluster_keywords[cluster_id] = get_top_keywords(cluster_sents, n=10)

    # sentences_df["topic_label"] = sentences_df["topic"].apply(
    #     lambda t: "Noise" if t == -1 else ", ".join(cluster_keywords.get(t, []))
    # )

    # Save
    sentences_df.to_csv(output_file, index=False)
    print(f"\n✅ Done! Saved clustered JTBD sentences with labels to {output_file}")


if __name__ == "__main__":
    import sys
    if len(sys.argv) < 3:
        print("Usage: python keep_jtbd_sentences.py input.csv output.csv")
    else:
        main(sys.argv[1], sys.argv[2])