"""Tests for GroupManager: schema, import, discovery funnel, hit_rate tracking."""

from __future__ import annotations

import csv
from pathlib import Path

import pytest

from agent_samochodowy.ingestor.groups import (
    DEAD_SCAN_THRESHOLD,
    ELIGIBLE_CATEGORIES,
    HARD_DEAD_SCANS,
    GroupManager,
)


@pytest.fixture
def gm(tmp_path):
    """In-memory GroupManager for tests."""
    db = str(tmp_path / "test.db")
    manager = GroupManager(db)
    yield manager
    manager.close()


def _insert_group(gm, group_id, nazwa=None, kategoria="czesci_tir", aktywna=1):
    """Helper to insert a single group."""
    nazwa = nazwa or f"Group {group_id}"
    gm.conn.execute("""
        INSERT INTO fb_groups (group_id, nazwa, url, kategoria, aktywna)
        VALUES (?, ?, ?, ?, ?)
    """, (group_id, nazwa, f"https://fb.com/groups/{group_id}", kategoria, aktywna))
    gm.conn.commit()


@pytest.fixture
def gm_with_groups(gm):
    """GroupManager pre-loaded with 5 czesci_tir groups."""
    for i in range(5):
        _insert_group(gm, f"g{i}")
    return gm


class TestSchema:
    def test_creates_tables(self, gm):
        cur = gm.conn.execute("SELECT name FROM sqlite_master WHERE type='table'")
        tables = {row[0] for row in cur.fetchall()}
        assert "fb_groups" in tables
        assert "group_scan_log" in tables

    def test_has_stats_columns(self, gm):
        cur = gm.conn.execute("PRAGMA table_info(fb_groups)")
        cols = {row[1] for row in cur.fetchall()}
        for col in ("scans_count", "hits_count", "queries_count",
                     "last_hit_at", "last_scraped_at", "group_state"):
            assert col in cols, f"Missing column: {col}"

    def test_group_state_default(self, gm):
        _insert_group(gm, "x1")
        cur = gm.conn.execute("SELECT group_state FROM fb_groups WHERE group_id='x1'")
        assert cur.fetchone()[0] == "untested"

    def test_migrate_idempotent(self, tmp_path):
        """Opening same DB twice should not fail (ALTER IF NOT EXISTS)."""
        db = str(tmp_path / "test.db")
        gm1 = GroupManager(db)
        gm1.close()
        gm2 = GroupManager(db)
        gm2.close()


class TestImportCsv:
    def test_import_and_count(self, gm, tmp_path):
        csv_path = tmp_path / "groups.csv"
        csv_path.write_text(
            "group_id,nazwa,url,kategoria,aktywna\n"
            "123,Test Group,https://fb.com/groups/123,czesci_tir,1\n"
            "456,Other Group,https://fb.com/groups/456,inne,0\n"
        )
        count = gm.import_csv(str(csv_path))
        assert count == 2
        stats = gm.stats()
        assert stats["total"] == 2
        assert stats["active"] == 1

    def test_upsert_updates_existing(self, gm, tmp_path):
        csv_path = tmp_path / "groups.csv"
        csv_path.write_text(
            "group_id,nazwa,url,kategoria,aktywna\n"
            "123,Old Name,https://fb.com/groups/123,czesci_tir,0\n"
        )
        gm.import_csv(str(csv_path))
        csv_path.write_text(
            "group_id,nazwa,url,kategoria,aktywna\n"
            "123,New Name,https://fb.com/groups/123,czesci_tir,1\n"
        )
        gm.import_csv(str(csv_path))
        cur = gm.conn.execute("SELECT nazwa, aktywna FROM fb_groups WHERE group_id='123'")
        row = cur.fetchone()
        assert row[0] == "New Name"
        assert row[1] == 1


class TestRecordScan:
    def test_updates_counters(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=5, hits=2)
        cur = gm.conn.execute(
            "SELECT scans_count, hits_count, queries_count, last_scraped_at, last_hit_at "
            "FROM fb_groups WHERE group_id='g0'"
        )
        row = cur.fetchone()
        assert row[0] == 1  # scans_count
        assert row[1] == 2  # hits_count
        assert row[2] == 5  # queries_count
        assert row[3] is not None  # last_scraped_at
        assert row[4] is not None  # last_hit_at (hits > 0)

    def test_no_hits_no_last_hit(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=3, hits=0)
        cur = gm.conn.execute("SELECT last_hit_at FROM fb_groups WHERE group_id='g0'")
        assert cur.fetchone()[0] is None

    def test_increments(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=2, hits=1)
        gm.record_scan("g0", queries=3, hits=0)
        cur = gm.conn.execute(
            "SELECT scans_count, hits_count, queries_count FROM fb_groups WHERE group_id='g0'"
        )
        row = cur.fetchone()
        assert row[0] == 2
        assert row[1] == 1
        assert row[2] == 5

    def test_scan_log_entries(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=2, hits=1)
        gm.record_scan("g0", queries=3, hits=2)
        cur = gm.conn.execute(
            "SELECT queries, hits FROM group_scan_log WHERE group_id='g0' ORDER BY id"
        )
        rows = cur.fetchall()
        assert len(rows) == 2
        assert rows[0][0] == 2 and rows[0][1] == 1
        assert rows[1][0] == 3 and rows[1][1] == 2


class TestRollingHitRate:
    def test_no_scans_returns_zero(self, gm_with_groups):
        assert gm_with_groups._rolling_hit_rate("g0") == 0.0

    def test_single_scan(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=5, hits=2)
        assert gm._rolling_hit_rate("g0") == 2.0  # 2 hits / 1 scan

    def test_window_limited(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=10, hits=10)
        gm.record_scan("g0", queries=10, hits=10)
        for _ in range(10):
            gm.record_scan("g0", queries=5, hits=0)
        assert gm._rolling_hit_rate("g0", window=10) == 0.0

    def test_hit_rate_reflects_recent(self, gm_with_groups):
        gm = gm_with_groups
        for _ in range(5):
            gm.record_scan("g0", queries=3, hits=0)
        for _ in range(5):
            gm.record_scan("g0", queries=3, hits=1)
        assert gm._rolling_hit_rate("g0") == 0.5


class TestWindowStats:
    def test_no_scans(self, gm_with_groups):
        scans, queries, hr = gm_with_groups._window_stats("g0")
        assert scans == 0
        assert queries == 0
        assert hr == 0.0

    def test_with_scans(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=3, hits=1)
        gm.record_scan("g0", queries=2, hits=0)
        scans, queries, hr = gm._window_stats("g0")
        assert scans == 2
        assert queries == 5
        assert hr == 0.5  # 1 hit / 2 scans


class TestComputeState:
    def test_untested_no_scans(self, gm_with_groups):
        assert gm_with_groups._compute_state("g0") == "untested"

    def test_untested_few_scans(self, gm_with_groups):
        gm = gm_with_groups
        for _ in range(DEAD_SCAN_THRESHOLD - 1):
            gm.record_scan("g0", queries=0, hits=0)
        assert gm._compute_state("g0") == "untested"

    def test_proven_with_queries(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=2, hits=1)
        assert gm._compute_state("g0") == "proven"

    def test_proven_even_with_many_empty_scans(self, gm_with_groups):
        """One query in window keeps group proven."""
        gm = gm_with_groups
        gm.record_scan("g0", queries=1, hits=0)
        for _ in range(DEAD_SCAN_THRESHOLD):
            gm.record_scan("g0", queries=0, hits=0)
        # query is still in the window (window=10, scans=7)
        assert gm._compute_state("g0") == "proven"

    def test_dead_threshold(self, gm_with_groups):
        gm = gm_with_groups
        for _ in range(DEAD_SCAN_THRESHOLD):
            gm.record_scan("g0", queries=0, hits=0)
        assert gm._compute_state("g0") == "dead"


class TestAutoCull:
    def test_marks_dead_inactive(self, gm):
        _insert_group(gm, "d1", aktywna=1)
        for _ in range(DEAD_SCAN_THRESHOLD):
            gm.record_scan("d1", queries=0, hits=0)
        result = gm.auto_cull()
        assert result["dead"] == 1
        assert result["culled_this_run"] == 1
        cur = gm.conn.execute("SELECT aktywna, group_state FROM fb_groups WHERE group_id='d1'")
        row = cur.fetchone()
        assert row[0] == 0
        assert row[1] == "dead"

    def test_marks_proven_active(self, gm):
        _insert_group(gm, "p1", aktywna=0)
        gm.record_scan("p1", queries=3, hits=1)
        result = gm.auto_cull()
        assert result["proven"] == 1
        cur = gm.conn.execute("SELECT aktywna, group_state FROM fb_groups WHERE group_id='p1'")
        row = cur.fetchone()
        assert row[0] == 1
        assert row[1] == "proven"

    def test_untested_unchanged(self, gm):
        _insert_group(gm, "u1", aktywna=0)
        gm.record_scan("u1", queries=0, hits=0)
        result = gm.auto_cull()
        assert result["untested"] == 1
        cur = gm.conn.execute("SELECT aktywna FROM fb_groups WHERE group_id='u1'")
        assert cur.fetchone()[0] == 0  # stays 0

    def test_ignores_non_eligible_categories(self, gm):
        _insert_group(gm, "t1", kategoria="transport", aktywna=1)
        for _ in range(DEAD_SCAN_THRESHOLD):
            gm.record_scan("t1", queries=0, hits=0)
        result = gm.auto_cull()
        # transport not in ELIGIBLE_CATEGORIES → not culled
        assert result.get("dead", 0) == 0
        cur = gm.conn.execute("SELECT aktywna FROM fb_groups WHERE group_id='t1'")
        assert cur.fetchone()[0] == 1  # unchanged

    def test_dead_revives_to_proven(self, gm):
        """A dead group that gets a query during retest → proven."""
        _insert_group(gm, "r1", aktywna=1)
        for _ in range(DEAD_SCAN_THRESHOLD):
            gm.record_scan("r1", queries=0, hits=0)
        gm.auto_cull()
        cur = gm.conn.execute("SELECT group_state FROM fb_groups WHERE group_id='r1'")
        assert cur.fetchone()[0] == "dead"

        # Simulate retest with a hit
        gm.record_scan("r1", queries=1, hits=1)
        gm.auto_cull()
        cur = gm.conn.execute("SELECT group_state, aktywna FROM fb_groups WHERE group_id='r1'")
        row = cur.fetchone()
        assert row[0] == "proven"
        assert row[1] == 1


class TestPickGroups:
    def test_returns_all_when_fewer_than_n(self, gm_with_groups):
        result = gm_with_groups.pick_groups(10)
        assert len(result) == 5

    def test_proven_gets_slots(self, gm):
        """Proven groups should get ~50% of slots."""
        # 3 proven, 3 untested
        for i in range(3):
            _insert_group(gm, f"p{i}")
            for _ in range(3):
                gm.record_scan(f"p{i}", queries=2, hits=1)
        for i in range(3):
            _insert_group(gm, f"u{i}")
        gm.auto_cull()

        # Pick 4 slots many times, proven should dominate
        proven_count = 0
        trials = 100
        for _ in range(trials):
            picked = gm.pick_groups(4)
            for g in picked:
                if g["group_state"] == "proven":
                    proven_count += 1
        # Expect proven to get at least 40% of total picks
        assert proven_count > trials * 4 * 0.3

    def test_untested_pool_excludes_inactive(self, gm):
        """Untested pool must only include active groups (aktywna=1)."""
        _insert_group(gm, "u1", kategoria="mechanika_porady", aktywna=0)
        _insert_group(gm, "u2", kategoria="mechanika_porady", aktywna=1)
        result = gm.pick_groups(3)
        ids = {g["group_id"] for g in result}
        assert "u1" not in ids
        assert "u2" in ids

    def test_untested_active_only(self, gm):
        """Only aktywna=1 untested groups enter the funnel."""
        for i in range(10):
            _insert_group(gm, f"inactive{i}", kategoria="czesci_tir", aktywna=0)
        _insert_group(gm, "active1", kategoria="czesci_tir", aktywna=1)
        result = gm.pick_groups(5)
        assert len(result) == 1
        assert result[0]["group_id"] == "active1"

    def test_retest_includes_dead(self, gm):
        """Dead groups should get ~15% retest slots."""
        # 2 proven, 2 dead, 0 untested
        for i in range(2):
            _insert_group(gm, f"p{i}")
            for _ in range(3):
                gm.record_scan(f"p{i}", queries=2, hits=1)
        for i in range(2):
            _insert_group(gm, f"d{i}")
            for _ in range(DEAD_SCAN_THRESHOLD):
                gm.record_scan(f"d{i}", queries=0, hits=0)
        gm.auto_cull()

        # Pick many times — dead should appear sometimes
        dead_seen = set()
        for _ in range(50):
            picked = gm.pick_groups(3)
            for g in picked:
                if g["group_state"] == "dead":
                    dead_seen.add(g["group_id"])
        assert len(dead_seen) > 0

    def test_empty_pool(self, gm):
        assert gm.pick_groups(3) == []

    def test_overflow_spills(self, gm):
        """When proven pool is empty, slots spill to untested."""
        for i in range(5):
            _insert_group(gm, f"u{i}")
        result = gm.pick_groups(3)
        assert len(result) == 3
        # All should be untested (no proven available)
        for g in result:
            assert g["group_state"] == "untested"

    def test_hit_rate_in_result(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=5, hits=2)
        result = gm.pick_groups(5)
        g0 = next(g for g in result if g["group_id"] == "g0")
        assert g0["hit_rate"] == 2.0

    def test_scans_in_window_in_result(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=0, hits=0)
        gm.record_scan("g0", queries=0, hits=0)
        result = gm.pick_groups(5)
        g0 = next(g for g in result if g["group_id"] == "g0")
        assert g0["scans_in_window"] == 2

    def test_non_eligible_excluded(self, gm):
        """Groups from non-eligible categories shouldn't be picked."""
        _insert_group(gm, "t1", kategoria="transport", aktywna=1)
        _insert_group(gm, "e1", kategoria="czesci_tir", aktywna=1)
        result = gm.pick_groups(5)
        ids = {g["group_id"] for g in result}
        assert "t1" not in ids
        assert "e1" in ids


class TestRanking:
    def test_ranking_order_by_state(self, gm):
        _insert_group(gm, "p1")
        gm.record_scan("p1", queries=3, hits=1)
        _insert_group(gm, "u1")
        _insert_group(gm, "d1")
        for _ in range(DEAD_SCAN_THRESHOLD):
            gm.record_scan("d1", queries=0, hits=0)
        gm.auto_cull()

        ranking = gm.ranking()
        states = [g["group_state"] for g in ranking]
        assert states == ["proven", "untested", "dead"]

    def test_ranking_includes_scans_in_window(self, gm):
        _insert_group(gm, "g1")
        gm.record_scan("g1", queries=0, hits=0)
        gm.record_scan("g1", queries=0, hits=0)
        ranking = gm.ranking()
        assert ranking[0]["scans_in_window"] == 2

    def test_ranking_shows_all_states(self, gm):
        """Ranking should show dead groups too (not just active)."""
        _insert_group(gm, "d1")
        for _ in range(DEAD_SCAN_THRESHOLD):
            gm.record_scan("d1", queries=0, hits=0)
        gm.auto_cull()

        ranking = gm.ranking()
        assert len(ranking) == 1
        assert ranking[0]["group_state"] == "dead"


class TestHardDead:
    def test_hard_dead_excluded_from_retest(self, gm):
        """Groups with >=HARD_DEAD_SCANS and 0 lifetime queries are never retested."""
        _insert_group(gm, "hd1", aktywna=1)
        # Simulate many scans with 0 queries
        for _ in range(HARD_DEAD_SCANS):
            gm.record_scan("hd1", queries=0, hits=0)
        gm.auto_cull()
        cur = gm.conn.execute("SELECT group_state FROM fb_groups WHERE group_id='hd1'")
        assert cur.fetchone()[0] == "dead"

        # Should NOT appear in retest pool
        dead_pool = gm._get_groups_by_state("dead", list(ELIGIBLE_CATEGORIES))
        ids = {g["group_id"] for g in dead_pool}
        assert "hd1" not in ids

    def test_soft_dead_still_retested(self, gm):
        """Groups with <HARD_DEAD_SCANS still appear in retest pool."""
        _insert_group(gm, "sd1", aktywna=1)
        for _ in range(DEAD_SCAN_THRESHOLD):
            gm.record_scan("sd1", queries=0, hits=0)
        gm.auto_cull()

        dead_pool = gm._get_groups_by_state("dead", list(ELIGIBLE_CATEGORIES))
        ids = {g["group_id"] for g in dead_pool}
        assert "sd1" in ids

    def test_hard_dead_with_past_queries_still_retested(self, gm):
        """Groups with many scans but some lifetime queries are not hard-dead."""
        _insert_group(gm, "mixed1", aktywna=1)
        gm.record_scan("mixed1", queries=1, hits=0)  # 1 query early on
        for _ in range(HARD_DEAD_SCANS):
            gm.record_scan("mixed1", queries=0, hits=0)
        gm.auto_cull()

        dead_pool = gm._get_groups_by_state("dead", list(ELIGIBLE_CATEGORIES))
        ids = {g["group_id"] for g in dead_pool}
        assert "mixed1" in ids  # queries_count=1 > 0, so not hard-dead

    def test_hard_dead_not_picked(self, gm):
        """pick_groups never returns hard-dead groups."""
        # One active untested + one hard-dead
        _insert_group(gm, "alive1", aktywna=1)
        _insert_group(gm, "hd1", aktywna=1)
        for _ in range(HARD_DEAD_SCANS):
            gm.record_scan("hd1", queries=0, hits=0)
        gm.auto_cull()

        for _ in range(20):
            picked = gm.pick_groups(3)
            for g in picked:
                assert g["group_id"] != "hd1"


class TestPostsReturned:
    def test_posts_returned_recorded(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=2, hits=1, posts_returned=15)
        cur = gm.conn.execute(
            "SELECT posts_returned FROM group_scan_log WHERE group_id='g0'"
        )
        assert cur.fetchone()[0] == 15

    def test_posts_returned_defaults_zero(self, gm_with_groups):
        gm = gm_with_groups
        gm.record_scan("g0", queries=0, hits=0)
        cur = gm.conn.execute(
            "SELECT posts_returned FROM group_scan_log WHERE group_id='g0'"
        )
        assert cur.fetchone()[0] == 0

    def test_scan_log_has_posts_returned_column(self, gm):
        cur = gm.conn.execute("PRAGMA table_info(group_scan_log)")
        cols = {row[1] for row in cur.fetchall()}
        assert "posts_returned" in cols


class TestStats:
    def test_by_state(self, gm):
        _insert_group(gm, "p1")
        gm.record_scan("p1", queries=2, hits=1)
        _insert_group(gm, "u1")
        _insert_group(gm, "d1")
        for _ in range(DEAD_SCAN_THRESHOLD):
            gm.record_scan("d1", queries=0, hits=0)
        gm.auto_cull()

        stats = gm.stats()
        assert stats["by_state"]["proven"] == 1
        assert stats["by_state"]["untested"] == 1
        assert stats["by_state"]["dead"] == 1
