import json
import os
from urllib.parse import urlparse
from openai import OpenAI
import sys

client = OpenAI(api_key=os.getenv("OPENAI_API_KEY"))

from collections import defaultdict
from urllib.parse import urlparse

def extract_subreddits_with_counts(data):
    subreddit_counts = defaultdict(int)
    for entry in data:
        for result in entry.get("results", []):
            url = result.get("url", "")
            if "reddit.com/r/" in url:
                try:
                    parts = urlparse(url).path.split("/")
                    subreddit = parts[2]
                    subreddit_counts[subreddit] += 1
                except IndexError:
                    continue
    return dict(sorted(subreddit_counts.items()))

def query_openai_for_relevance(subreddits):
    subreddit_counts_str = "\n".join([f"{sub}: {count}" for sub, count in subreddits.items()])
                                      
    prompt = f"""
You are helping filter Reddit communities for a project focused on budgetbackers wallet app.

Each subreddit below includes a number of posts that matched this query:

("budgetbakers" OR "wallet by budgetbakers" OR "wallet app") site:reddit.com inurl:comments

Each subreddit is followed by the number of posts it matched in the dataset.

Your task: Return a Python list of subreddit names that are likely to contain **relevant personal experiences or struggles**, based on the subreddit *topic* and *number of posts*. Exclude gaming communities.

Return only a valid Python list of subreddit names. No explanations. Just the list.

Subreddits and their post counts:

{subreddit_counts_str}
"""
    response = client.chat.completions.create(
        model="gpt-4",
        messages=[
            {"role": "system", "content": "You are a helpful assistant that classifies subreddit relevance."},
            {"role": "user", "content": prompt}
        ],
        temperature=0
    )

    try:
        relevant_list_str = response.choices[0].message.content.strip()
        relevant_subreddits = eval(relevant_list_str)
        return set(relevant_subreddits)
    except Exception as e:
        print("Error parsing OpenAI response:", e)
        return set()

def filter_data_by_subreddits(data, relevant_subreddits):
    filtered_data = []
    for entry in data:
        filtered_results = []
        for result in entry.get("results", []):
            url = result.get("url", "")
            if "reddit.com/r/" in url:
                try:
                    parts = urlparse(url).path.split("/")
                    subreddit = parts[2]
                    if subreddit in relevant_subreddits:
                        filtered_results.append(result)
                except IndexError:
                    continue
        if filtered_results:
            entry["results"] = filtered_results
            filtered_data.append(entry)
    return filtered_data

# Main
if __name__ == "__main__":
    input_file = sys.argv[1]
    with open(input_file, "r", encoding="utf-8") as f:
        data = json.load(f)

    subreddits = extract_subreddits_with_counts(data)
    print(f"Found {len(subreddits)} unique subreddits.")

    print(subreddits)

    relevant_subreddits = query_openai_for_relevance(subreddits)
    print(f"{len(relevant_subreddits)} subreddits deemed relevant by OpenAI.")
    print(relevant_subreddits)

    filtered_data = filter_data_by_subreddits(data, relevant_subreddits)


    with open(input_file, "w", encoding="utf-8") as f:
        json.dump(filtered_data, f, indent=2, ensure_ascii=False)

    print("Filtered data saved to filtered_data.json.")
