import pandas as pd
from tqdm import tqdm

import sys

# ---------------- CONFIG ----------------
JTBD_FILE = sys.argv[1]
COMMENTS_FILE = sys.argv[2]
OUTPUT_FILE = "similar_comments.csv"
TOP_N = 5
MODEL_NAME = "all-MiniLM-L6-v2"
# ----------------------------------------

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

# ---------------- CONFIG ----------------
OUTPUT_FILE = "similar_comments_openai.csv"
COMMENTS_EMBED_FILE = "comment_embeddings.pkl"
MODEL = "text-embedding-3-large"
TOP_N = 5
# ----------------------------------------

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

def get_embedding(text: str) -> list:
    """Generate an embedding for a single text using OpenAI."""
    text = text.replace("\n", " ")
    response = openai.embeddings.create(
        model=MODEL,
        input=[text]
    )
    return response.data[0].embedding

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

def load_or_create_embeddings(comments):
    """Load embeddings from disk if available, otherwise generate and save."""
    if Path(COMMENTS_EMBED_FILE).exists():
        print(f"🔹 Loading cached embeddings from {COMMENTS_EMBED_FILE} ...")
        with open(COMMENTS_EMBED_FILE, "rb") as f:
            return pickle.load(f)

    print("🔹 Generating embeddings for comments...")
    comment_embeddings = []
    for c in tqdm(comments, desc="Encoding comments"):
        emb = get_embedding(c)
        comment_embeddings.append(emb)

    with open(COMMENTS_EMBED_FILE, "wb") as f:
        pickle.dump(comment_embeddings, f)

    return comment_embeddings

def main():
    print("🔹 Loading data...")
    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 ["comment", "situation", "struggle", "outcome"]), \
        "jtbd.csv must have columns: comment, situation, struggle, outcome"

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

    results = []

    print("🔹 Processing JTBDs...")
    for idx, row in tqdm(jtbd_df.iterrows(), total=len(jtbd_df)):
        jtbd_id = idx
        for field in ["situation", "struggle", "outcome"]:
            query_text = str(row[field])
            if not query_text.strip():
                continue

            query_emb = get_embedding(query_text)
            sims = [cosine_similarity(query_emb, e) for e in comment_embeddings]

            top_idx = np.argsort(sims)[::-1][:TOP_N]
            for i in top_idx:
                results.append({
                    "jtbd_id": jtbd_id,
                    "jtbd_field": field,
                    "original_text": query_text,
                    "similar_comment": comments[i],
                    "similarity_score": float(sims[i])
                })

    print(f"💾 Saving results to {OUTPUT_FILE} ...")
    pd.DataFrame(results).to_csv(OUTPUT_FILE, index=False)
    print("✅ Done! File saved as", OUTPUT_FILE)

if __name__ == "__main__":
    main()
