import sys
import os
import json
import csv
import time
from urllib.parse import urlparse

from selenium import webdriver
from selenium.webdriver.common.by import By
from selenium.webdriver.chrome.options import Options

# ✅ Input
if len(sys.argv) < 2:
    print("Usage: python scrape_reddit_post_details_json.py <input_json_file>")
    sys.exit(1)

input_json = sys.argv[1]

# ✅ Load posts from JSON (expects an object with a "results" list)
with open(input_json, "r", encoding="utf-8") as f:
    data = json.load(f)

posts = data[0]["results"]

# ✅ Derive comments CSV name
# Prefer r_<subreddit>_comments.csv if we can detect the subreddit from the first valid URL
def derive_comments_csv(input_file):
    name = input_file.split('.')[0]
    return name + '_comments.csv'

def is_subreddit_root(url):
    """
    Detect if the URL is in the form:
      https://www.reddit.com/r/<subreddit>/ (with optional trailing slash)
    and no further path—i.e., not a specific post.
    """
    try:
        parsed = urlparse(url)
        path = parsed.path.strip("/")  # e.g. "r/gallbladders" or "r/gallbladders/comments/..."
        parts = path.split("/")
        return len(parts) == 2 and parts[0].lower() == "r"
    except:
        return False

base = os.path.splitext(os.path.basename(input_json))[0]
comments_csv = derive_comments_csv(input_json)

# ✅ Setup browser — DO NOT CHANGE THESE OPTIONS
options = Options()
options.add_argument("--headless")
options.add_argument("--disable-gpu")
options.add_argument("--no-sandbox")
options.add_argument("--window-size=1920x1080")
options.add_argument("user-agent=Mozilla/5.0")
driver = webdriver.Chrome(options=options)

# ✅ Helpers
def save_json_state():
    with open(input_json, "w", encoding="utf-8") as f:
        json.dump(data, f, ensure_ascii=False, indent=2)

def append_comments(rows, file_path):
    if not rows:
        return
    file_exists = os.path.isfile(file_path)
    with open(file_path, "a", newline="", encoding="utf-8") as f:
        writer = csv.DictWriter(f, fieldnames=["post_url", "author", "comment", "date"])
        if not file_exists:
            writer.writeheader()
        writer.writerows(rows)

def get_post_content(driver):
    # Try several robust containers for new Reddit
    for sel in [
        "[data-test-id='post-content']",
        "div[data-click-id='text']",
        "shreddit-post",
        "article",
    ]:
        try:
            el = driver.find_element(By.CSS_SELECTOR, sel)
            txt = el.text.strip()
            if txt:
                return txt
        except:
            pass
    # Fallback to headers
    for sel in ["h1", "h2", "h3"]:
        try:
            el = driver.find_element(By.CSS_SELECTOR, sel)
            txt = el.text.strip()
            if txt:
                return txt
        except:
            pass
    return ""

def get_post_datetime(driver):
    # Try common locations for the post timestamp
    for sel in [
        "faceplate-timeago time[datetime]",
        "time[datetime]"
    ]:
        try:
            el = driver.find_element(By.CSS_SELECTOR, sel)
            return el.get_attribute("datetime")
        except:
            pass
    return None

# 🔁 Scrape loop (per-post isolation; never aborts early)
for idx, post in enumerate(posts):
    url = post.get("url")
    if not url:
        print(f"[{idx+1}/{len(posts)}] Skipping (no 'url' field).")
        # Still mark as processed to avoid infinite retries if desired:
        # post["scraped_comments"] = "True"; save_json_state()
        continue

    # Skip if it's subreddit root
    if is_subreddit_root(url):
        print(f"[{idx+1}/{len(posts)}] Ignoring subreddit root URL: {url}")
        continue

    if str(post.get("scraped_comments", "")).lower() == "true":
        print(f"[{idx+1}/{len(posts)}] Skipping (already scraped): {url}")
        continue

    print(f"[{idx+1}/{len(posts)}] Scraping: {url}")

    try:
        driver.get(url)
        time.sleep(4)

        # Post content
        post_content = ""
        try:
            post_content = get_post_content(driver)
        except Exception as e:
            print(f"  ⚠️ Post content error: {e}")

        # Post datetime
        post_datetime = ""
        try:
            post_datetime = get_post_datetime(driver)
        except Exception as e:
            print(f"  ⚠️ Post datetime error: {e}")


        # Comments
        comments = []
        try:
            comment_blocks = driver.find_elements(By.CSS_SELECTOR, "shreddit-comment")
            for block in comment_blocks:
                parsed_comment = {}
                try:
                    author = block.get_attribute("author") or "unknown"
                    text = block.text.strip()
                    if text:
                        parsed_comment = {
                            "post_url": url,
                            "author": author,
                            "comment": text,
                            "date": ""
                        }

                    # grab comment datetime
                    date = None
                    try:
                        t = block.find_element(By.CSS_SELECTOR, "a time[datetime]")
                        date = t.get_attribute("datetime")
                    except:
                        try:
                            menu = block.find_element(
                                By.CSS_SELECTOR,
                                "shreddit-overflow-menu[comment-created-timestamp]"
                            )
                            date = menu.get_attribute("comment-created-timestamp")
                        except:
                            pass
                        
                    if date:
                        parsed_comment["date"] = date

                    if len(parsed_comment) > 0:
                        comments.append(parsed_comment)
                except:
                    continue
        except Exception as e:
            print(f"  ⚠️ Comments error: {e}")

        # Save comments to separate CSV
        append_comments(comments, comments_csv)
        print(f"  💾 Saved {len(comments)} comments -> {comments_csv}")

        # Update JSON (content + mark as done)
        if post_content:
            post["post_content"] = post_content

        if post_datetime:
            post["datetime"] = post_datetime
        post["scraped_comments"] = "True"

    except Exception as e:
        print(f"  ❌ Unexpected error for this URL: {e}")

    # Always save after each post for resumability
    save_json_state()
    time.sleep(2)

driver.quit()
print("\n✅ DONE: Updated JSON with post content & flags; comments saved to separate CSV.")
