"""Facebook groups table: import, discovery funnel, stats tracking.

Groups flow through a three-state funnel:
  untested → proven (has queries) or dead (>=N scans, 0 queries)

Selection uses three tiers:
  ~50% proven (exploit), ~35% untested (screen), ~15% retest (probe dead).
"""

from __future__ import annotations

import csv
import random
import sqlite3
from datetime import datetime
from pathlib import Path


# Rolling window for hit_rate and state calculations
HIT_RATE_WINDOW = 10

# Scans with 0 queries before a group is marked dead
DEAD_SCAN_THRESHOLD = 6

# Lifetime scans with 0 queries ever → permanently dead, never retest
HARD_DEAD_SCANS = 50

# Categories eligible for the discovery funnel
ELIGIBLE_CATEGORIES = ("czesci_tir", "czesci_inne", "felgi_opony", "mechanika_porady")

# Three-tier slot allocation
TIER_PROVEN = 0.50
TIER_UNTESTED = 0.35
TIER_RETEST = 0.15


class GroupManager:
    """Manages the fb_groups table and discovery funnel in catalog.db."""

    def __init__(self, db_path: str) -> None:
        Path(db_path).parent.mkdir(parents=True, exist_ok=True)
        self.conn = sqlite3.connect(db_path)
        self.conn.row_factory = sqlite3.Row
        self._init_schema()

    # ------------------------------------------------------------------
    # Schema
    # ------------------------------------------------------------------

    def _init_schema(self) -> None:
        self.conn.execute("""
            CREATE TABLE IF NOT EXISTS fb_groups (
                group_id TEXT PRIMARY KEY,
                nazwa TEXT NOT NULL,
                url TEXT NOT NULL,
                kategoria TEXT NOT NULL DEFAULT '',
                aktywna INTEGER NOT NULL DEFAULT 0,
                last_scraped_at TEXT,
                last_hit_at TEXT,
                scans_count INTEGER NOT NULL DEFAULT 0,
                hits_count INTEGER NOT NULL DEFAULT 0,
                queries_count INTEGER NOT NULL DEFAULT 0
            )
        """)
        # Per-scan log for rolling window hit_rate
        self.conn.execute("""
            CREATE TABLE IF NOT EXISTS group_scan_log (
                id INTEGER PRIMARY KEY AUTOINCREMENT,
                group_id TEXT NOT NULL,
                scanned_at TEXT NOT NULL,
                queries INTEGER NOT NULL DEFAULT 0,
                hits INTEGER NOT NULL DEFAULT 0,
                posts_returned INTEGER NOT NULL DEFAULT 0
            )
        """)
        self.conn.execute(
            "CREATE INDEX IF NOT EXISTS idx_scan_log_group ON group_scan_log(group_id)"
        )
        # Migrate: add columns if missing (idempotent)
        for col, defn in [
            ("last_hit_at", "TEXT"),
            ("scans_count", "INTEGER NOT NULL DEFAULT 0"),
            ("hits_count", "INTEGER NOT NULL DEFAULT 0"),
            ("queries_count", "INTEGER NOT NULL DEFAULT 0"),
            ("group_state", "TEXT NOT NULL DEFAULT 'untested'"),
        ]:
            try:
                self.conn.execute(f"ALTER TABLE fb_groups ADD COLUMN {col} {defn}")
            except sqlite3.OperationalError:
                pass  # column already exists
        # Migrate scan_log: add posts_returned if missing
        try:
            self.conn.execute(
                "ALTER TABLE group_scan_log ADD COLUMN posts_returned INTEGER NOT NULL DEFAULT 0"
            )
        except sqlite3.OperationalError:
            pass
        self.conn.commit()

    # ------------------------------------------------------------------
    # Import
    # ------------------------------------------------------------------

    def import_csv(self, csv_path: str) -> int:
        """UPSERT groups from CSV. Returns count of rows imported."""
        count = 0
        with open(csv_path, newline="", encoding="utf-8") as f:
            reader = csv.DictReader(f)
            for row in reader:
                self.conn.execute("""
                    INSERT INTO fb_groups (group_id, nazwa, url, kategoria, aktywna)
                    VALUES (?, ?, ?, ?, ?)
                    ON CONFLICT(group_id) DO UPDATE SET
                        nazwa=excluded.nazwa,
                        url=excluded.url,
                        kategoria=excluded.kategoria,
                        aktywna=excluded.aktywna
                """, (
                    row["group_id"].strip(),
                    row["nazwa"].strip(),
                    row["url"].strip(),
                    row.get("kategoria", "").strip(),
                    int(row.get("aktywna", 0)),
                ))
                count += 1
        self.conn.commit()
        return count

    # ------------------------------------------------------------------
    # State computation
    # ------------------------------------------------------------------

    def _window_stats(
        self, group_id: str, window: int = HIT_RATE_WINDOW,
    ) -> tuple[int, int, float]:
        """Return (scans_in_window, queries_in_window, hit_rate) for a group."""
        cur = self.conn.execute("""
            SELECT hits, queries FROM group_scan_log
            WHERE group_id = ?
            ORDER BY scanned_at DESC
            LIMIT ?
        """, (group_id, window))
        rows = cur.fetchall()
        if not rows:
            return 0, 0, 0.0
        scans = len(rows)
        total_queries = sum(r["queries"] for r in rows)
        total_hits = sum(r["hits"] for r in rows)
        hit_rate = total_hits / scans
        return scans, total_queries, hit_rate

    def _compute_state(self, group_id: str) -> str:
        """Compute group_state from scan history in the rolling window."""
        scans, queries, _ = self._window_stats(group_id)
        if queries > 0:
            return "proven"
        if scans >= DEAD_SCAN_THRESHOLD:
            return "dead"
        return "untested"

    def auto_cull(self, categories: tuple[str, ...] | None = None) -> dict[str, int]:
        """Recompute group_state for all groups in eligible categories.

        - proven → aktywna=1
        - dead   → aktywna=0  (removed from rotation, kept in DB)
        - untested → aktywna unchanged

        Returns {state: count, "culled_this_run": N}.
        """
        cats = list(categories or ELIGIBLE_CATEGORIES)
        placeholders = ",".join("?" for _ in cats)
        cur = self.conn.execute(f"""
            SELECT group_id, group_state, aktywna FROM fb_groups
            WHERE kategoria IN ({placeholders})
        """, cats)

        counts: dict[str, int] = {"untested": 0, "proven": 0, "dead": 0, "culled_this_run": 0}

        for row in cur.fetchall():
            gid = row["group_id"]
            old_state = row["group_state"]
            new_state = self._compute_state(gid)
            counts[new_state] += 1

            sets: list[str] = []
            vals: list[object] = []

            if new_state != old_state:
                sets.append("group_state = ?")
                vals.append(new_state)
            if new_state == "proven" and row["aktywna"] != 1:
                sets.append("aktywna = ?")
                vals.append(1)
            if new_state == "dead" and row["aktywna"] != 0:
                sets.append("aktywna = ?")
                vals.append(0)
                counts["culled_this_run"] += 1

            if sets:
                vals.append(gid)
                self.conn.execute(
                    f"UPDATE fb_groups SET {', '.join(sets)} WHERE group_id = ?",
                    vals,
                )

        self.conn.commit()
        return counts

    # ------------------------------------------------------------------
    # Three-tier group selection
    # ------------------------------------------------------------------

    def pick_groups(
        self, n: int, kategoria: str | list[str] | None = None,
    ) -> list[dict]:
        """Pick n groups using three-tier discovery funnel.

        Tiers:
        - ~50% proven  (exploit — weighted by hit_rate)
        - ~35% untested (screen — oldest-first from eligible categories)
        - ~15% retest   (probe dead groups for revival)

        Empty-tier slots spill to tiers with remaining capacity.
        """
        cats = self._resolve_categories(kategoria)

        proven = self._get_groups_by_state("proven", cats)
        untested = self._get_groups_by_state("untested", cats)
        dead = self._get_groups_by_state("dead", cats)

        all_available = proven + untested + dead
        if not all_available:
            return []
        if len(all_available) <= n:
            return all_available

        # Allocate slots: guarantee 1 per non-empty tier, then distribute
        # remaining by ratio.
        tiers = [
            ("proven", proven, TIER_PROVEN),
            ("untested", untested, TIER_UNTESTED),
            ("retest", dead, TIER_RETEST),
        ]
        non_empty = [(name, pool, ratio) for name, pool, ratio in tiers if pool]
        pools = {"proven": proven, "untested": untested, "retest": dead}

        if n <= len(non_empty):
            # Fewer slots than tiers — give 1 each in priority order
            slots = {name: 1 if i < n else 0
                     for i, (name, _, _) in enumerate(non_empty)}
        else:
            remaining = n - len(non_empty)
            ratio_sum = sum(r for _, _, r in non_empty) or 1.0
            slots = {}
            for name, _, ratio in non_empty:
                extra = round(remaining * ratio / ratio_sum)
                slots[name] = 1 + extra
            # Fix rounding to match n exactly
            diff = n - sum(slots.values())
            for name, _, _ in non_empty:
                if diff == 0:
                    break
                if diff > 0:
                    slots[name] += 1
                    diff -= 1
                elif slots[name] > 1:
                    slots[name] -= 1
                    diff += 1

        # Ensure we have entries for all tier names
        for name in ("proven", "untested", "retest"):
            slots.setdefault(name, 0)

        # Clamp to actual pool sizes, spill overflow
        overflow = 0
        for tier in list(slots):
            avail = len(pools[tier])
            if slots[tier] > avail:
                overflow += slots[tier] - avail
                slots[tier] = avail
        for tier in ("untested", "proven", "retest"):
            if overflow <= 0:
                break
            spare = len(pools[tier]) - slots[tier]
            take = min(overflow, spare)
            slots[tier] += take
            overflow -= take

        picked: list[dict] = []

        # Proven: weighted random by hit_rate
        if proven and slots["proven"] > 0:
            proven.sort(key=lambda g: (-g["hit_rate"], g["last_scraped_at"] or ""))
            picked.extend(self._weighted_sample(proven, slots["proven"]))

        # Untested: oldest-first (never scraped → first)
        if untested and slots["untested"] > 0:
            untested.sort(
                key=lambda g: (
                    g["last_scraped_at"] is not None,
                    g["last_scraped_at"] or "",
                ),
            )
            picked.extend(untested[: slots["untested"]])

        # Retest: oldest last_scraped_at (least recently checked)
        if dead and slots["retest"] > 0:
            dead.sort(key=lambda g: (g["last_scraped_at"] or ""))
            picked.extend(dead[: slots["retest"]])

        # Deduplicate and trim
        seen = set()
        result = []
        for g in picked:
            if g["group_id"] not in seen:
                seen.add(g["group_id"])
                result.append(g)
        return result[:n]

    # ------------------------------------------------------------------
    # Internals
    # ------------------------------------------------------------------

    def _resolve_categories(self, kategoria: str | list[str] | None) -> list[str]:
        if kategoria is None:
            return list(ELIGIBLE_CATEGORIES)
        if isinstance(kategoria, str):
            return [kategoria]
        return list(kategoria)

    def _get_groups_by_state(self, state: str, categories: list[str]) -> list[dict]:
        """Fetch groups by state from given categories, enriched with window stats.

        For proven/untested: only active groups (aktywna=1).
        For dead (retest): excludes hard-dead groups (>=HARD_DEAD_SCANS with 0 lifetime queries).
        """
        placeholders = ",".join("?" for _ in categories)

        if state in ("proven", "untested"):
            cur = self.conn.execute(f"""
                SELECT group_id, nazwa, url, kategoria, aktywna,
                       last_scraped_at, last_hit_at,
                       scans_count, hits_count, queries_count, group_state
                FROM fb_groups
                WHERE kategoria IN ({placeholders})
                  AND group_state = ?
                  AND aktywna = 1
            """, [*categories, state])
        else:
            # dead / retest: skip hard-dead groups
            cur = self.conn.execute(f"""
                SELECT group_id, nazwa, url, kategoria, aktywna,
                       last_scraped_at, last_hit_at,
                       scans_count, hits_count, queries_count, group_state
                FROM fb_groups
                WHERE kategoria IN ({placeholders})
                  AND group_state = ?
                  AND NOT (scans_count >= ? AND queries_count = 0)
            """, [*categories, state, HARD_DEAD_SCANS])

        groups = []
        for row in cur.fetchall():
            g = dict(row)
            scans, queries, hit_rate = self._window_stats(g["group_id"])
            g["hit_rate"] = hit_rate
            g["scans_in_window"] = scans
            g["queries_in_window"] = queries
            groups.append(g)
        return groups

    def _get_active_groups(self, kategoria: str | list[str]) -> list[dict]:
        """Fetch all active groups with rolling hit_rate (legacy compat)."""
        if isinstance(kategoria, str):
            kategoria = [kategoria]
        placeholders = ",".join("?" for _ in kategoria)
        cur = self.conn.execute(f"""
            SELECT group_id, nazwa, url, kategoria, aktywna,
                   last_scraped_at, last_hit_at,
                   scans_count, hits_count, queries_count, group_state
            FROM fb_groups
            WHERE kategoria IN ({placeholders}) AND aktywna = 1
        """, kategoria)
        groups = []
        for row in cur.fetchall():
            g = dict(row)
            scans, queries, hit_rate = self._window_stats(g["group_id"])
            g["hit_rate"] = hit_rate
            g["scans_in_window"] = scans
            g["queries_in_window"] = queries
            groups.append(g)
        return groups

    def _weighted_sample(self, groups: list[dict], n: int) -> list[dict]:
        """Weighted random sample: higher hit_rate → higher chance."""
        weights = [g["hit_rate"] + 0.05 for g in groups]
        n = min(n, len(groups))
        picked: list[dict] = []
        available = list(range(len(groups)))
        available_weights = list(weights)
        for _ in range(n):
            if not available:
                break
            [idx] = random.choices(available, weights=available_weights, k=1)
            picked.append(groups[idx])
            pos = available.index(idx)
            available.pop(pos)
            available_weights.pop(pos)
        return picked

    def _rolling_hit_rate(self, group_id: str, window: int = HIT_RATE_WINDOW) -> float:
        """Compute hit_rate from last `window` scans."""
        _, _, hit_rate = self._window_stats(group_id, window)
        return hit_rate

    # ------------------------------------------------------------------
    # Post-scan recording
    # ------------------------------------------------------------------

    def record_scan(
        self, group_id: str, queries: int, hits: int, posts_returned: int = 0,
    ) -> None:
        """Record scan results and update group counters."""
        now = datetime.now().isoformat()

        # Log entry for rolling window
        self.conn.execute("""
            INSERT INTO group_scan_log (group_id, scanned_at, queries, hits, posts_returned)
            VALUES (?, ?, ?, ?, ?)
        """, (group_id, now, queries, hits, posts_returned))

        # Update aggregate counters
        self.conn.execute("""
            UPDATE fb_groups SET
                last_scraped_at = ?,
                scans_count = scans_count + 1,
                hits_count = hits_count + ?,
                queries_count = queries_count + ?
            WHERE group_id = ?
        """, (now, hits, queries, group_id))

        if hits > 0:
            self.conn.execute(
                "UPDATE fb_groups SET last_hit_at = ? WHERE group_id = ?",
                (now, group_id),
            )

        self.conn.commit()

    def mark_scraped(self, group_id: str) -> None:
        """Legacy: update last_scraped_at only (use record_scan for full tracking)."""
        self.record_scan(group_id, queries=0, hits=0)

    # ------------------------------------------------------------------
    # Diagnostics
    # ------------------------------------------------------------------

    def ranking(self, kategoria: str | list[str] | None = None) -> list[dict]:
        """Return all groups from eligible categories, ranked by state then hit_rate."""
        cats = self._resolve_categories(kategoria)
        placeholders = ",".join("?" for _ in cats)
        cur = self.conn.execute(f"""
            SELECT group_id, nazwa, url, kategoria, aktywna,
                   last_scraped_at, last_hit_at,
                   scans_count, hits_count, queries_count, group_state
            FROM fb_groups
            WHERE kategoria IN ({placeholders})
        """, cats)
        state_order = {"proven": 0, "untested": 1, "dead": 2}
        groups = []
        for row in cur.fetchall():
            g = dict(row)
            scans, queries, hit_rate = self._window_stats(g["group_id"])
            g["hit_rate"] = hit_rate
            g["scans_in_window"] = scans
            g["queries_in_window"] = queries
            groups.append(g)
        groups.sort(
            key=lambda g: (
                state_order.get(g["group_state"], 9),
                -g["hit_rate"],
                -g["scans_count"],
            ),
        )
        return groups

    def stats(self) -> dict:
        """Return summary stats including state breakdown."""
        cur = self.conn.execute("SELECT COUNT(*) FROM fb_groups")
        total = cur.fetchone()[0]
        cur = self.conn.execute("SELECT COUNT(*) FROM fb_groups WHERE aktywna = 1")
        active = cur.fetchone()[0]
        cur = self.conn.execute(
            "SELECT COUNT(*) FROM fb_groups WHERE kategoria = 'czesci_tir' AND aktywna = 1"
        )
        active_tir = cur.fetchone()[0]

        # State counts for eligible categories
        placeholders = ",".join("?" for _ in ELIGIBLE_CATEGORIES)
        cur = self.conn.execute(f"""
            SELECT group_state, COUNT(*) FROM fb_groups
            WHERE kategoria IN ({placeholders})
            GROUP BY group_state
        """, list(ELIGIBLE_CATEGORIES))
        by_state = {row[0]: row[1] for row in cur.fetchall()}

        cur = self.conn.execute(
            "SELECT kategoria, COUNT(*) FROM fb_groups GROUP BY kategoria ORDER BY COUNT(*) DESC"
        )
        by_cat = {row[0]: row[1] for row in cur.fetchall()}
        return {
            "total": total,
            "active": active,
            "active_czesci_tir": active_tir,
            "by_category": by_cat,
            "by_state": by_state,
        }

    def close(self) -> None:
        self.conn.close()
