import os
import pandas as pd
import numpy as np
from tqdm import tqdm
import openai
import pickle
from pathlib import Path

import sys

# ---------------- CONFIG ----------------
JTBD_FILE = sys.argv[1]
COMMENTS_FILE = sys.argv[2]
OUTPUT_FILE = "jtbd_semantic_matches.csv"
COMMENTS_EMBED_CACHE = "comment_embeddings.pkl"
MODEL = "text-embedding-3-large"
TOP_N = 30
# ----------------------------------------

openai.api_key = os.getenv("OPENAI_API_KEY")

def get_embedding_batch(texts):
    """Return embeddings for a batch of texts."""
    response = openai.embeddings.create(model=MODEL, input=texts)
    return [d.embedding for d in response.data]

def cosine_similarity(a, b):
    a = np.array(a)
    b = np.array(b)
    return np.dot(a, b) / (np.linalg.norm(a) * np.linalg.norm(b))

def load_or_create_comment_embeddings(comments):
    """Load cached comment embeddings or generate them once."""
    if Path(COMMENTS_EMBED_CACHE).exists():
        print(f"🔹 Loading cached embeddings from {COMMENTS_EMBED_CACHE}")
        with open(COMMENTS_EMBED_CACHE, "rb") as f:
            return pickle.load(f)

    print("🔹 Generating embeddings for all comments (batched)...")
    all_embeddings = []
    BATCH = 100
    for i in tqdm(range(0, len(comments), BATCH)):
        batch = comments[i:i+BATCH]
        all_embeddings.extend(get_embedding_batch(batch))

    with open(COMMENTS_EMBED_CACHE, "wb") as f:
        pickle.dump(all_embeddings, f)
    return all_embeddings

def get_jtbd_vector(row):
    """Compute a semantic signature vector for a JTBD (average of situation+struggle+outcome)."""
    parts = [str(row.get(col, "")).strip() for col in ["situation", "struggle", "outcome"] if str(row.get(col, "")).strip()]
    if not parts:
        return None
    text = " ".join(parts)
    emb = get_embedding_batch([text])[0]
    return emb

def main():
    print("🔹 Loading CSVs...")
    jtbd_df = pd.read_csv(JTBD_FILE)
    comments_df = pd.read_csv(COMMENTS_FILE)

    assert "comment" in comments_df.columns, "comments.csv must have a 'comment' column"
    assert all(c in jtbd_df.columns for c in ["situation", "struggle", "outcome"]), \
        "jtbd.csv must include columns: situation, struggle, outcome"

    comments = comments_df["comment"].astype(str).tolist()
    comment_embeddings = load_or_create_comment_embeddings(comments)

    results = []
    print("🔹 Processing JTBDs...")
    for idx, row in tqdm(jtbd_df.iterrows(), total=len(jtbd_df)):
        jtbd_vector = get_jtbd_vector(row)
        if jtbd_vector is None:
            continue

        sims = [cosine_similarity(jtbd_vector, e) for e in comment_embeddings]
        top_idx = np.argsort(sims)[::-1][:TOP_N]

        for rank, i in enumerate(top_idx, 1):
            results.append({
                "jtbd_id": idx,
                "jtbd_text": f"Situation: {row['situation']} | Struggle: {row['struggle']} | Outcome: {row['outcome']}",
                "similar_comment": comments[i],
                "similarity_score": float(sims[i]),
                "rank": rank
            })

    out_df = pd.DataFrame(results)
    out_df.to_csv(OUTPUT_FILE, index=False)
    print(f"✅ Done! Saved {len(out_df)} rows to {OUTPUT_FILE}")

if __name__ == "__main__":
    main()
