# -----------------------
# 1. Setup
# -----------------------
import pandas as pd
import re
import spacy
from sentence_transformers import SentenceTransformer
from sklearn.cluster import KMeans
import matplotlib.pyplot as plt
from sklearn.decomposition import PCA
import sys

# Load spaCy model (English)
nlp = spacy.load("en_core_web_sm")

# Situation trigger words
situation_keywords = ["when", "while", "after", "before", "if", "because", "since", "as"]

# -----------------------
# 2. Load comments
# -----------------------
# Example: comments stored in a CSV with a column called "comment"
df = pd.read_csv(sys.argv[1])
comments = df["comment"].dropna().tolist()

# -----------------------
# 3. Extract situation phrases
# -----------------------
def extract_situations(text):
    if not isinstance(text, str):
        return None  # skip NaN or non-string rows
    
    doc = nlp(text)
    situations = []
    
    # Dependency parsing
    for token in doc:
        if token.dep_ == "mark" and token.text.lower() in situation_keywords:
            phrase = " ".join([t.text for t in token.subtree])
            situations.append(phrase)
    
    # Regex backup
    regex_hits = re.findall(
        r"\b(?:when|while|after|before|if|because|since|as)\b[^.?!]+",
        text,
        re.IGNORECASE
    )
    situations.extend(regex_hits)
    
    situations = list(set(situations))
    
    # Print each situation found
    for s in situations:
        print(f"Found situation: {s}")
    
    return situations if situations else None


# Apply extraction
df["situations"] = df["comment"].apply(extract_situations)
df = df.explode("situations").dropna(subset=["situations"]).reset_index(drop=True)

print("Extracted situations:")
print(df[["comment", "situations"]].head(10))

# -----------------------
# 4. Embed and cluster situations
# -----------------------
model = SentenceTransformer("all-MiniLM-L6-v2")
situation_texts = df["situations"].tolist()
embeddings = model.encode(situation_texts, show_progress_bar=True)

# KMeans clustering
num_clusters = 8   # tune based on dataset size
kmeans = KMeans(n_clusters=num_clusters, random_state=42)
df["cluster"] = kmeans.fit_predict(embeddings)

# -----------------------
# 5. Review clusters
# -----------------------
for c in range(num_clusters):
    examples = df[df["cluster"] == c]["situations"].head(5).tolist()
    print(f"\nCluster {c} examples:")
    for e in examples:
        print(" -", e)

# -----------------------
# 6. Optional: visualize clusters
# -----------------------
pca = PCA(n_components=2)
reduced = pca.fit_transform(embeddings)

plt.figure(figsize=(10, 6))
scatter = plt.scatter(reduced[:,0], reduced[:,1], c=df["cluster"], cmap="tab10", alpha=0.7)
plt.legend(*scatter.legend_elements(), title="Cluster")
plt.title("Situation Clusters (PCA projection)")
plt.show()

# -----------------------
# 7. Save results
# -----------------------
df.to_csv("reddit_situations_clustered.csv", index=False)
