import argparse
import sqlite3
from pathlib import Path

BASE_DIR = Path(__file__).resolve().parent
DB_DIR = BASE_DIR / "incoming_uploads" / "sqlite_output"

SOURCES = [
    ("posts.db", "posts_data", "Posts"),
    ("whatsapp.db", "whatsapp_data", "WhatsApp"),
    ("xiaohongshu.db", "xiaohongshu_data", "Xiaohongshu"),
    ("wechat.db", "wechat_data", "Wechat"),
]

CANDIDATE_AUTHORS = [
    "author", "author_name", "username", "user_name", "user", "sender",
    "nickname", "account", "screen_name", "reply_author", "comment_author", "name", "from"
]

CATEGORY_NAME = "Location, Transport & Environment"
ENV_SUBCATS = {"view", "nearby environment", "pest", "noise", "feng shui", "fengshui"}


def quote_ident(s):
    s = s.replace('"', '""')
    return f'"{s}"'


def classify_row(cat_raw, sub_raw, env_cat):
    cat_clean = cat_raw.strip().lower()
    sub_clean = sub_raw.strip().lower()
    if env_cat.strip().lower() == cat_clean:
        if sub_clean in ENV_SUBCATS:
            return "Environment & View"
        return "Location & Transport"
    return cat_raw.strip()


def get_columns(con, table):
    rows = con.execute(f"PRAGMA table_info({quote_ident(table)})").fetchall()
    if not rows:
        raise ValueError(f"Table {table} not found or has no columns.")
    return [r[1] for r in rows]


def resolve_author_column(table_cols, override, source_label):
    cols_lower = {c.lower(): c for c in table_cols}
    if override:
        override_lower = override.lower()
        if override_lower in cols_lower:
            return cols_lower[override_lower]
        raise ValueError(
            f"Specified author column '{override}' not found in table. "
            f"Available columns: {', '.join(table_cols)}"
        )
    for cand in CANDIDATE_AUTHORS:
        if cand in cols_lower:
            return cols_lower[cand]
    raise ValueError(
        f"Could not auto-detect author column for source '{source_label}'. "
        f"Available columns: {', '.join(table_cols)}"
    )


def main():
    parser = argparse.ArgumentParser(description="SQLite Sentiment Report")
    parser.add_argument(
        "--author-column",
        action="append",
        default=[],
        help="Override author column in format SOURCE=COLUMN (e.g., posts.db=sender)"
    )
    args = parser.parse_args()

    overrides = {}
    for item in args.author_column:
        if "=" not in item:
            raise ValueError(f"Invalid override format '{item}', expected SOURCE=COLUMN")
        src, col = item.split("=", 1)
        overrides[src.strip()] = col.strip()

    for db_name, table_name, label in SOURCES:
        db_path = DB_DIR / db_name
        if not db_path.exists():
            print(f"[{label}] Database file not found: {db_path}")
            print("-" * 80)
            continue

        try:
            uri = f"file:{db_path.as_posix()}?mode=ro"
            con = sqlite3.connect(uri, uri=True)
            con.row_factory = sqlite3.Row
            table_cols = get_columns(con, table_name)

            required = ["cat", "sub_cat", "sentiment"]
            cols_lower = {c.lower(): c for c in table_cols}
            missing = [r for r in required if r not in cols_lower]
            if missing:
                raise ValueError(f"Missing required columns: {', '.join(missing)}. "
                                 f"Available columns: {', '.join(table_cols)}")

            col_cat = cols_lower["cat"]
            col_sub = cols_lower["sub_cat"]
            col_sent = cols_lower["sentiment"]
            
            ov = overrides.get(db_name)
            col_author = resolve_author_column(table_cols, ov, label)

            q_cat = quote_ident(col_cat)
            q_sub = quote_ident(col_sub)
            q_sent = quote_ident(col_sent)
            q_auth = quote_ident(col_author)

            query = f"""
            SELECT {q_cat} as cat, {q_sub} as sub_cat, {q_sent} as sent, {q_auth} as auth
            FROM {quote_ident(table_name)}
            """
            rows = con.execute(query).fetchall()
            con.close()
        except Exception as e:
            print(f"[{label}] ERROR: {e}")
            print("-" * 80)
            continue

        stats = {}
        for row in rows:
            cat_raw = row["cat"] or ""
            sub_raw = row["sub_cat"] or ""
            sent_raw = (row["sent"] or "").strip().lower()
            auth_raw = (row["auth"] or "").strip().lower()
            
            mapped_cat = classify_row(cat_raw, sub_raw, CATEGORY_NAME)

            if mapped_cat not in stats:
                stats[mapped_cat] = {
                    "total": 0,
                    "authors": set(),
                    "neg": 0, "neg_authors": set(),
                    "pos": 0, "pos_authors": set(),
                    "neu": 0, "neu_authors": set(),
                }

            bucket = stats[mapped_cat]
            bucket["total"] += 1
            if auth_raw:
                bucket["authors"].add(auth_raw)

            if sent_raw.startswith("neg"):
                bucket["neg"] += 1
                if auth_raw:
                    bucket["neg_authors"].add(auth_raw)
            elif sent_raw.startswith("pos"):
                bucket["pos"] += 1
                if auth_raw:
                    bucket["pos_authors"].add(auth_raw)
            elif sent_raw.startswith("neu"):
                bucket["neu"] += 1
                if auth_raw:
                    bucket["neu_authors"].add(auth_raw)

        print(f"Source: {label}")
        print("-" * 80)
        
        cat_w = 30
        print(f"{'Category':<{cat_w}} | {'Total':>6} | {'Authors':>7} | {'Neg':>4} | {'NegAuth':>7} | {'Pos':>4} | {'PosAuth':>7} | {'Neu':>4} | {'NeuAuth':>7}")
        
        for cat in sorted(stats.keys()):
            b = stats[cat]
            print(
                f"{cat:<{cat_w}} | "
                f"{b['total']:>6} | "
                f"{len(b['authors']):>7} | "
                f"{b['neg']:>4} | "
                f"{len(b['neg_authors']):>7} | "
                f"{b['pos']:>4} | "
                f"{len(b['pos_authors']):>7} | "
                f"{b['neu']:>4} | "
                f"{len(b['neu_authors']):>7}"
            )
        print("-" * 80)


if __name__ == "__main__":
    main()