"""
Sweep Slack messages for a user and store them into SQLite.
Run tools/index_slack.py separately to push stored data into the embedding engine.

Usage:
    python tools/fetch_slack.py                  # incremental (new messages only)
    python tools/fetch_slack.py --full           # full re-sweep; resumes if interrupted
    python tools/fetch_slack.py --full --reset   # full re-sweep, truly start over

Required env vars (set in .env or environment):
    SLACK_USER_TOKEN        xoxp-... user token (needs channels/im/groups read+history scopes)
    SLACK_USER_ID           Slack user ID to sweep (e.g. U01ABC123)
    SLACK_SWEEP_REPO_ID     repo_id (default: slack-<user_id>)
    SLACK_SWEEP_SNAPSHOT_ID snapshot_id (default: snap-<YYYY-MM-DD>)
    SLACK_SWEEP_DAYS        How many days back to fetch on first run (default: 365)
    SLACK_SQLITE_DB         Path to SQLite DB (default: tools/slack_data.db)
"""

import os
import re
import sys
import sqlite3
import time
import logging
from datetime import datetime, timedelta, timezone
from pathlib import Path

sys.path.insert(0, str(Path(__file__).parent.parent))

from dotenv import load_dotenv
load_dotenv(Path(__file__).parent.parent / ".env")

from slack_sdk import WebClient
from slack_sdk.errors import SlackApiError

logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
log = logging.getLogger("fetch_slack")

# ── config ────────────────────────────────────────────────────────────────────

SLACK_USER_TOKEN = os.environ["SLACK_USER_TOKEN"]
SLACK_USER_ID    = os.environ["SLACK_USER_ID"]
TODAY            = datetime.now(timezone.utc).strftime("%Y-%m-%d")
REPO_ID          = os.environ.get("SLACK_SWEEP_REPO_ID",     f"slack-{SLACK_USER_ID}")
SNAPSHOT_ID      = os.environ.get("SLACK_SWEEP_SNAPSHOT_ID", f"snap-{TODAY}")
# Default to 365 days on first full sweep; incremental runs use saved per-channel ts
SWEEP_DAYS       = int(os.environ.get("SLACK_SWEEP_DAYS", "365"))
FALLBACK_OLDEST  = str((datetime.now(timezone.utc) - timedelta(days=SWEEP_DAYS)).timestamp())
SQLITE_DB        = os.environ.get("SLACK_SQLITE_DB", str(Path(__file__).parent / "slack_data.db"))

client = WebClient(token=SLACK_USER_TOKEN)

# ── User name cache ───────────────────────────────────────────────────────────

_user_cache = {}

def get_user_name(user_id: str) -> str:
    """Resolve Slack user ID to real name, with caching."""
    if user_id in _user_cache:
        return _user_cache[user_id]
    
    for attempt in range(5):
        try:
            resp = client.users_info(user=user_id)
            user = resp.get("user", {})
            profile = user.get("profile", {})
            
            # Try multiple name fields in order of preference
            name = (
                profile.get("real_name_normalized") or
                profile.get("real_name") or
                profile.get("display_name_normalized") or
                profile.get("display_name") or
                user.get("real_name") or
                user.get("name") or
                user_id
            )
            
            # For bots, prefix with bot indicator
            if user.get("is_bot") and name and name != user_id:
                name = f"[Bot] {name}"
            
            _user_cache[user_id] = name
            time.sleep(0.2)  # rate limit
            return name
        except SlackApiError as exc:
            if exc.response.get("error") == "ratelimited":
                wait = int(exc.response.headers.get("Retry-After", 30))
                log.warning("Rate limited on users_info — waiting %ds (attempt %d/5)", wait, attempt + 1)
                time.sleep(wait)
                continue
            log.warning("Failed to resolve user %s: %s", user_id, exc.response.get("error"))
            _user_cache[user_id] = user_id
            return user_id
    
    # Max retries exceeded
    _user_cache[user_id] = user_id
    return user_id

# ── Slack helpers ─────────────────────────────────────────────────────────────

def _paginate(method, **kwargs):
    """Yield all pages from a cursor-paginated Slack API method."""
    cursor = None
    while True:
        for attempt in range(5):
            try:
                resp = method(**kwargs, cursor=cursor, limit=200)
                break
            except SlackApiError as exc:
                if exc.response.get("error") == "ratelimited":
                    wait = int(exc.response.headers.get("Retry-After", 30))
                    log.warning("Rate limited — waiting %ds (attempt %d/5)", wait, attempt + 1)
                    time.sleep(wait)
                else:
                    raise
        else:
            log.error("Max retries exceeded for %s", method.__name__)
            return
        yield from resp.get("channels") or resp.get("messages") or []
        cursor = (resp.get("response_metadata") or {}).get("next_cursor")
        if not cursor:
            break
        time.sleep(1.0)  # respect Slack rate limits (Tier 3: ~50 req/min)


def list_conversations():
    """Return all conversations the authed user is a member of."""
    log.info("Listing conversations for user %s", SLACK_USER_ID)
    convos = list(_paginate(
        client.conversations_list,
        types="public_channel,private_channel,im,mpim",
        exclude_archived=True,
    ))
    # keep only channels the user is actually a member of
    return [c for c in convos if c.get("is_member") or c.get("is_im") or c.get("is_mpim")]


def fetch_messages(channel_id: str, oldest_ts: str) -> list[dict]:
    """Return all messages in channel since oldest_ts."""
    messages = []
    cursor = None
    while True:
        try:
            resp = client.conversations_history(
                channel=channel_id,
                oldest=oldest_ts,
                limit=200,
                cursor=cursor,
            )
        except SlackApiError as exc:
            if exc.response.get("error") == "ratelimited":
                wait = int(exc.response.headers.get("Retry-After", 10))
                log.warning("Rate limited — waiting %ds", wait)
                time.sleep(wait)
                continue
            log.warning("Skipping channel %s: %s", channel_id, exc.response.get("error"))
            return []

        for msg in resp.get("messages", []):
            if msg.get("type") != "message" or msg.get("subtype"):
                continue
            messages.append(msg)
            # fetch full thread for every threaded message
            if msg.get("reply_count"):
                for attempt in range(5):
                    try:
                        thread_resp = client.conversations_replies(
                            channel=channel_id,
                            ts=msg["ts"],
                            oldest=oldest_ts,
                            limit=200,
                        )
                        # skip first item (it's the parent, already added)
                        messages.extend(thread_resp.get("messages", [])[1:])
                        time.sleep(0.5)
                        break
                    except SlackApiError as exc:
                        if exc.response.get("error") == "ratelimited":
                            wait = int(exc.response.headers.get("Retry-After", 30))
                            log.warning("Rate limited on replies — waiting %ds (attempt %d/5)", wait, attempt + 1)
                            time.sleep(wait)
                            continue
                        break  # other errors, skip thread

        cursor = (resp.get("response_metadata") or {}).get("next_cursor")
        if not cursor:
            break
        time.sleep(1.0)  # increased delay between history pages

    return messages


# ── document builder ──────────────────────────────────────────────────────────

def build_doc(channel: dict, messages: list[dict]) -> dict | None:
    if not messages:
        return None

    channel_name = channel.get("name") or channel.get("id")
    lines = [f"# Slack channel: #{channel_name}  ({TODAY})\n"]

    user_ids = set()
    for m in messages:
        if m.get("user"):
            user_ids.add(m["user"])
        for match in re.finditer(r'<@(U[A-Z0-9]+)>', m.get("text") or ""):
            user_ids.add(match.group(1))

    log.info("  Resolving %d user names for #%s", len(user_ids), channel_name)
    for uid in user_ids:
        get_user_name(uid)

    for m in sorted(messages, key=lambda x: float(x.get("ts", 0))):
        ts = datetime.fromtimestamp(float(m["ts"]), tz=timezone.utc).strftime("%Y-%m-%d %H:%M")
        user_id = m.get("user", "unknown")
        user_name = get_user_name(user_id) if user_id != "unknown" else "unknown"
        text = (m.get("text") or "").strip()
        if text:
            text = re.sub(r'<@(U[A-Z0-9]+)>', lambda x: f"@{get_user_name(x.group(1))}", text)
            lines.append(f"[{ts}] {user_name}: {text}")

    return {
        "repo_id": REPO_ID,
        "doc_path": f"slack/{channel_name}/{TODAY}.md",
        "content": "\n".join(lines),
        "commit_hash": SNAPSHOT_ID,
    }


# ── SQLite storage ────────────────────────────────────────────────────────────

def init_db() -> sqlite3.Connection:
    con = sqlite3.connect(SQLITE_DB)
    con.execute("""
        CREATE TABLE IF NOT EXISTS slack_documents (
            id           INTEGER PRIMARY KEY AUTOINCREMENT,
            repo_id      TEXT NOT NULL,
            channel_id   TEXT NOT NULL,
            channel_name TEXT NOT NULL,
            doc_path     TEXT NOT NULL,
            snapshot_id  TEXT NOT NULL,
            sweep_date   TEXT NOT NULL,
            content      TEXT NOT NULL,
            latest_ts    TEXT,
            swept_at     TEXT NOT NULL,
            indexed      INTEGER NOT NULL DEFAULT 0,
            indexed_at   TEXT,
            -- one row per channel per calendar day; same-day re-runs update in place
            UNIQUE(repo_id, channel_id, sweep_date)
        )
    """)
    # per-channel sweep state for incremental runs
    con.execute("""
        CREATE TABLE IF NOT EXISTS slack_sweep_state (
            channel_id   TEXT PRIMARY KEY,
            channel_name TEXT,
            last_ts      TEXT NOT NULL,
            last_swept   TEXT NOT NULL
        )
    """)
    # per-channel progress so an interrupted sweep can resume without re-fetching done channels
    con.execute("""
        CREATE TABLE IF NOT EXISTS slack_sweep_progress (
            sweep_date TEXT NOT NULL,
            channel_id TEXT NOT NULL,
            done       INTEGER NOT NULL DEFAULT 0,
            saved_at   TEXT,
            PRIMARY KEY (sweep_date, channel_id)
        )
    """)
    # stores the start date of an in-progress full sweep so it survives day boundaries
    con.execute("""
        CREATE TABLE IF NOT EXISTS slack_config (
            key   TEXT PRIMARY KEY,
            value TEXT NOT NULL
        )
    """)
    con.commit()
    return con


def _is_done(con: sqlite3.Connection, channel_id: str, full: bool = False, full_sweep_start: str | None = None) -> bool:
    if full and full_sweep_start:
        # check across all dates since the full sweep began, not just today
        row = con.execute(
            "SELECT done FROM slack_sweep_progress WHERE channel_id=? AND done=1 AND sweep_date>=?",
            (channel_id, full_sweep_start),
        ).fetchone()
    else:
        row = con.execute(
            "SELECT done FROM slack_sweep_progress WHERE sweep_date=? AND channel_id=?",
            (TODAY, channel_id),
        ).fetchone()
    return bool(row and row[0])


def _mark_done(con: sqlite3.Connection, channel_id: str) -> None:
    con.execute("""
        INSERT INTO slack_sweep_progress (sweep_date, channel_id, done, saved_at)
        VALUES (?, ?, 1, ?)
        ON CONFLICT(sweep_date, channel_id) DO UPDATE SET done=1, saved_at=excluded.saved_at
    """, (TODAY, channel_id, datetime.now(timezone.utc).isoformat()))
    con.commit()


def get_oldest_ts(con: sqlite3.Connection, channel_id: str, full: bool) -> str:
    """Return oldest_ts for this channel: saved last_ts or global fallback."""
    if not full:
        row = con.execute(
            "SELECT last_ts FROM slack_sweep_state WHERE channel_id = ?", (channel_id,)
        ).fetchone()
        if row:
            return row[0]
    return FALLBACK_OLDEST


def save_to_db(con: sqlite3.Connection, channel_id: str, cname: str, doc: dict, latest_ts: str | None) -> None:
    con.execute("""
        INSERT INTO slack_documents
            (repo_id, channel_id, channel_name, doc_path, snapshot_id, sweep_date, content, latest_ts, swept_at, indexed)
        VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, 0)
        ON CONFLICT(repo_id, channel_id, sweep_date) DO UPDATE SET
            content    = excluded.content,
            latest_ts  = excluded.latest_ts,
            swept_at   = excluded.swept_at,
            indexed    = 0,
            indexed_at = NULL
    """, (
        doc["repo_id"], channel_id, cname,
        doc["doc_path"], doc["commit_hash"], TODAY,
        doc["content"], latest_ts, datetime.now(timezone.utc).isoformat(),
    ))
    if latest_ts:
        con.execute("""
            INSERT INTO slack_sweep_state (channel_id, channel_name, last_ts, last_swept)
            VALUES (?, ?, ?, ?)
            ON CONFLICT(channel_id) DO UPDATE SET
                last_ts    = excluded.last_ts,
                last_swept = excluded.last_swept
        """, (channel_id, cname, latest_ts, datetime.now(timezone.utc).isoformat()))
    con.commit()


# ── main ──────────────────────────────────────────────────────────────────────

def run(full: bool = False, reset: bool = False):
    con = init_db()

    if reset:
        con.execute("DELETE FROM slack_sweep_progress")
        con.execute("DELETE FROM slack_sweep_state")
        con.execute("DELETE FROM slack_config WHERE key='full_sweep_start'")
        con.commit()
        log.info("All progress and sweep state cleared — starting from scratch")

    # for --full runs, persist the sweep start date so progress survives midnight rollovers
    if full:
        row = con.execute("SELECT value FROM slack_config WHERE key='full_sweep_start'").fetchone()
        if row and not reset:
            full_sweep_start = row[0]
            log.info("Resuming full sweep started on %s", full_sweep_start)
        else:
            full_sweep_start = TODAY
            con.execute(
                "INSERT INTO slack_config(key,value) VALUES('full_sweep_start',?) "
                "ON CONFLICT(key) DO UPDATE SET value=excluded.value",
                (full_sweep_start,),
            )
            con.commit()
    else:
        full_sweep_start = None

    log.info("SQLite DB: %s  |  full=%s  reset=%s  |  user=%s  |  repo=%s",
             SQLITE_DB, full, reset, SLACK_USER_ID, REPO_ID)

    convos = list_conversations()
    log.info("Found %d conversations to sweep", len(convos))

    if full and full_sweep_start:
        done_count = con.execute(
            "SELECT COUNT(*) FROM slack_sweep_progress WHERE done=1 AND sweep_date>=?",
            (full_sweep_start,),
        ).fetchone()[0]
    else:
        done_count = con.execute(
            "SELECT COUNT(*) FROM slack_sweep_progress WHERE sweep_date=? AND done=1", (TODAY,)
        ).fetchone()[0]
    if done_count:
        log.info("Resuming — %d channels already done, skipping them", done_count)

    saved = 0
    for idx, channel in enumerate(convos, 1):
        cid   = channel["id"]
        cname = channel.get("name") or cid

        if _is_done(con, cid, full=full, full_sweep_start=full_sweep_start):
            log.info("[%d/%d] Skipping #%s (already swept)", idx, len(convos), cname)
            continue

        oldest_ts = get_oldest_ts(con, cid, full)
        log.info("[%d/%d] Fetching #%s (since %s)", idx, len(convos), cname, oldest_ts[:10])

        messages = fetch_messages(cid, oldest_ts)
        if not messages:
            _mark_done(con, cid)  # no messages — still checkpoint so we don't retry
            continue

        # track latest message ts for next incremental run
        latest_ts = max((m.get("ts", "0") for m in messages), default=None)

        doc = build_doc(channel, messages)
        if not doc:
            if latest_ts:
                con.execute("""
                    INSERT INTO slack_sweep_state (channel_id, channel_name, last_ts, last_swept)
                    VALUES (?, ?, ?, ?)
                    ON CONFLICT(channel_id) DO UPDATE SET
                        last_ts=excluded.last_ts, last_swept=excluded.last_swept
                """, (cid, cname, latest_ts, datetime.now(timezone.utc).isoformat()))
                con.commit()
            _mark_done(con, cid)
            log.info("  No relevant messages in #%s — skipping", cname)
            continue

        save_to_db(con, cid, cname, doc, latest_ts)
        _mark_done(con, cid)
        log.info("  Saved #%s  (messages=%d  latest_ts=%s)", cname, len(messages), (latest_ts or "")[:10])
        saved += 1
        time.sleep(1.0)  # proactive pacing — Slack Tier 3 is ~50 req/min across all calls

    if full:
        # full sweep finished — clear start marker so next --full begins fresh
        con.execute("DELETE FROM slack_config WHERE key='full_sweep_start'")
        con.commit()
    con.close()
    log.info("Done. %d documents saved to %s", saved, SQLITE_DB)
    log.info("Run: py tools/index_slack.py  to push to embedding engine")


if __name__ == "__main__":
    full_sweep = "--full" in sys.argv
    reset_all  = "--reset" in sys.argv
    run(full=full_sweep, reset=reset_all)
