#!/usr/bin/env python3
"""Scheduled pipeline run: scrape groups, extract, match, write to dgred.

Designed for systemd timer (every 2h). Includes:
- Three-tier group selection (50% proven / 35% untested / 15% retest)
- Freshness filter (POST_MAX_AGE_HOURS)
- Circuit breaker (max leads + max errors per cycle)
- Per-cycle stats logged to cycle_log table
- Daily report generation
"""

from __future__ import annotations

import json
import logging
import re
import sqlite3
import sys
import time
from collections import Counter
from datetime import datetime, timedelta
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))

import httpx

from agent_samochodowy.account_guard import AccountGuard
from agent_samochodowy.billing.ledger import BillingLedger
from agent_samochodowy.config import Settings
from agent_samochodowy.dgred.client import DgredClient
from agent_samochodowy.dgred.writer import LeadWriter
from agent_samochodowy.extraction import extract
from agent_samochodowy.ingestor.groups import GroupManager
from agent_samochodowy.matcher.engine import match
from agent_samochodowy.matcher.index import CatalogIndex
from agent_samochodowy.models import Post
from agent_samochodowy.status import map_status

LOG_DIR = Path("/var/log/agent-samochodowy")
LOG_DIR.mkdir(parents=True, exist_ok=True)

logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s %(levelname)s %(name)s %(message)s",
    handlers=[
        logging.StreamHandler(),
        logging.FileHandler(LOG_DIR / "pipeline.log"),
    ],
)
logger = logging.getLogger("scheduled_run")
logging.getLogger("httpx").setLevel(logging.WARNING)
logging.getLogger("openai").setLevel(logging.WARNING)


def _mask_token(text: str) -> str:
    """Mask API tokens in log messages to prevent credential leaks."""
    return re.sub(r"(apify_api_)[A-Za-z0-9]+", r"\1***", text)
logging.getLogger("httpcore").setLevel(logging.WARNING)

APIFY_BASE = "https://api.apify.com/v2"
RESULTS_LIMIT = 20


def _use_auth_scraper(settings: Settings) -> bool:
    """Return True if the authenticated (cookies-based) scraper should be used."""
    return bool(settings.fb_cookies and settings.fb_cookies.strip())


# ------------------------------------------------------------------
# Cycle log (persistent stats per run)
# ------------------------------------------------------------------

def _ensure_cycle_log(db_path: str) -> None:
    conn = sqlite3.connect(db_path)
    conn.execute("""
        CREATE TABLE IF NOT EXISTS cycle_log (
            id INTEGER PRIMARY KEY AUTOINCREMENT,
            started_at TEXT NOT NULL,
            finished_at TEXT,
            groups_scanned INTEGER DEFAULT 0,
            groups_failed INTEGER DEFAULT 0,
            posts_fetched INTEGER DEFAULT 0,
            posts_too_old INTEGER DEFAULT 0,
            posts_analyzed INTEGER DEFAULT 0,
            queries_found INTEGER DEFAULT 0,
            leads_written INTEGER DEFAULT 0,
            leads_matched INTEGER DEFAULT 0,
            leads_review INTEGER DEFAULT 0,
            leads_no_match INTEGER DEFAULT 0,
            duplicates INTEGER DEFAULT 0,
            errors INTEGER DEFAULT 0,
            halted_reason TEXT
        )
    """)
    conn.commit()
    conn.close()


def _log_cycle(db_path: str, data: dict) -> int:
    """Insert cycle log row and return the new cycle id."""
    conn = sqlite3.connect(db_path)
    cols = ", ".join(data.keys())
    placeholders = ", ".join("?" for _ in data)
    cur = conn.execute(f"INSERT INTO cycle_log ({cols}) VALUES ({placeholders})", list(data.values()))
    cycle_id = cur.lastrowid
    conn.commit()
    conn.close()
    return cycle_id or 0


# ------------------------------------------------------------------
# Apify scraping
# ------------------------------------------------------------------

def _build_actor_input(settings: Settings, group_url: str) -> tuple[str, dict]:
    """Return (actor_id, input_json) for the appropriate scraper."""
    if _use_auth_scraper(settings):
        cookies = json.loads(settings.fb_cookies)
        input_json: dict = {
            "startUrls": [{"url": group_url}],
            "maxPosts": RESULTS_LIMIT,
            "viewOption": "CHRONOLOGICAL",
            "maxConcurrency": 1,
            "cookies": cookies,
        }
        if settings.fb_proxy_url:
            input_json["proxyUrl"] = settings.fb_proxy_url
        return settings.apify_actor_id_auth, input_json
    # Legacy public scraper (fallback)
    return settings.apify_actor_id, {
        "startUrls": [{"url": group_url}],
        "resultsLimit": RESULTS_LIMIT,
        "sortBy": "New posts",
    }


def scrape_group(settings: Settings, group_url: str) -> tuple[list[dict], str]:
    """Scrape a single group. Returns (items, run_status)."""
    actor_id, input_json = _build_actor_input(settings, group_url)
    token = settings.apify_token

    with httpx.Client(timeout=180) as client:
        try:
            resp = client.post(
                f"{APIFY_BASE}/acts/{actor_id}/runs",
                params={"token": token},
                json=input_json,
            )
            resp.raise_for_status()
        except Exception as e:
            logger.warning("Apify start failed for %s: %s", group_url, _mask_token(str(e)))
            return [], "START_FAILED"

        run_data = resp.json().get("data", {})
        run_id = run_data.get("id")
        dataset_id = run_data.get("defaultDatasetId")
        if not run_id:
            return [], "NO_RUN_ID"

        status = "UNKNOWN"
        for _ in range(90):
            time.sleep(2)
            try:
                sr = client.get(f"{APIFY_BASE}/actor-runs/{run_id}", params={"token": token})
                status = sr.json().get("data", {}).get("status", "UNKNOWN")
                if status in ("SUCCEEDED", "FAILED", "ABORTED", "TIMED-OUT"):
                    break
            except Exception:
                continue

        if status != "SUCCEEDED":
            logger.warning("Apify run %s ended: %s", run_id, status)
            return [], status

        try:
            items_resp = client.get(
                f"{APIFY_BASE}/datasets/{dataset_id}/items",
                params={"token": token, "format": "json"},
            )
            items_resp.raise_for_status()
            return items_resp.json(), status
        except Exception as e:
            logger.warning("Dataset fetch failed: %s", _mask_token(str(e)))
            return [], "DATASET_FAILED"


def map_item(item: dict, group_name: str) -> Post | None:
    text = item.get("text") or item.get("message") or ""
    if not text.strip():
        return None
    post_id = str(item.get("postId") or item.get("id") or item.get("url", ""))
    if not post_id:
        return None

    timestamp = item.get("timestamp") or item.get("time") or item.get("date")
    if timestamp and isinstance(timestamp, str):
        for fmt in ("%Y-%m-%dT%H:%M:%S.%fZ", "%Y-%m-%dT%H:%M:%SZ",
                    "%Y-%m-%dT%H:%M:%S", "%Y-%m-%d %H:%M:%S"):
            try:
                timestamp = datetime.strptime(timestamp, fmt)
                break
            except ValueError:
                continue
        else:
            timestamp = datetime.now()
    elif not timestamp:
        timestamp = datetime.now()

    author = (item.get("authorName")
              or (item.get("user", {}).get("name") if isinstance(item.get("user"), dict) else None)
              or "Unknown")
    profile_url = item.get("authorProfileUrl") or item.get("profileUrl") or ""
    post_url = item.get("postUrl") or item.get("url") or ""

    try:
        return Post(
            post_id=post_id, group=group_name, author_name=author,
            author_profile_url=profile_url, text=text,
            timestamp=timestamp, post_url=post_url,
        )
    except Exception:
        return None


# ------------------------------------------------------------------
# Main
# ------------------------------------------------------------------

def main() -> int:
    settings = Settings()
    started_at = datetime.now().isoformat()

    _ensure_cycle_log(settings.db_path)

    use_auth = _use_auth_scraper(settings)
    guard = AccountGuard(settings)

    # --- Pre-flight checks (auth scraper only) ---
    if use_auth:
        if guard.is_suspended():
            logger.warning("=== Cycle SKIPPED: account suspended ===")
            guard.close()
            return 0

        if not guard.check_daily_cap():
            logger.warning(
                "=== Cycle SKIPPED: daily scan cap reached (%d/%d) ===",
                guard.daily_count(), settings.max_daily_scans,
            )
            guard.close()
            return 0

    logger.info(
        "=== Scheduled run started (dry_run=%s, billing=%s, scraper=%s) ===",
        settings.dry_run, settings.billing_mode,
        "auth" if use_auth else "legacy",
    )

    # Reserve cycle_id for deterministic billing references
    _ensure_cycle_log(settings.db_path)
    conn_tmp = sqlite3.connect(settings.db_path)
    cur_tmp = conn_tmp.execute(
        "INSERT INTO cycle_log (started_at) VALUES (?)", (started_at,)
    )
    cycle_id = cur_tmp.lastrowid or 0
    conn_tmp.commit()
    conn_tmp.close()

    billing = BillingLedger(settings, cycle_id=cycle_id)

    # --- Token balance pre-flight (live mode) ---
    if billing.is_live:
        if not billing.check_balance_sufficient(min_tokens=settings.tokens_per_scan):
            # Tokens exhausted — send alarm if newly exhausted
            billing.send_tokens_alarm(settings, exhausted=True, balance=0)
            logger.warning("=== Cycle SKIPPED: tokens exhausted ===")
            billing.close()
            guard.close()
            return 0
        # Balance restored after previous exhaustion — send restoration alarm
        # (check_balance_sufficient already auto-cleared the flag and logged)

    # Warmup: dynamic group count
    warmup_n = guard.warmup_group_count()
    n_groups = warmup_n if warmup_n is not None else settings.n_groups_per_cycle
    if warmup_n is not None:
        logger.info("Warmup active: %d groups (until %s)", warmup_n, settings.warmup_until)

    gm = GroupManager(settings.db_path)
    groups = gm.pick_groups(n_groups)
    logger.info("Selected %d groups", len(groups))

    # Scrape (with billing per scan)
    all_posts: list[Post] = []
    groups_ok = 0
    groups_failed = 0
    groups_skipped_billing = 0
    group_post_map: dict[str, list[Post]] = {}  # group_id → posts for that group

    group_posts_count: dict[str, int] = {}  # group_id → raw posts returned by Apify

    for g in groups:
        if use_auth:
            # Daily cap check per group (auth mode only)
            if not guard.check_daily_cap():
                logger.warning("Daily scan cap reached mid-cycle, stopping scraping")
                break

        raw_items, run_status = scrape_group(settings, g["url"])

        if use_auth:
            # Account guard: check run status
            guard.check_run_failure(run_status)
            if guard.is_suspended():
                logger.error("Account suspended during scraping — aborting cycle")
                break

        # Filter error items (works for both legacy and auth)
        if use_auth:
            items = guard.check_items(g["nazwa"], raw_items) if raw_items else []
            if guard.is_suspended():
                logger.error("Account suspended (threat in items) — aborting cycle")
                break
            guard.record_group_result(len(items))
            guard.increment_daily_count()
        else:
            # Legacy mode: simple error filtering, no guard checks
            items = [i for i in raw_items if not i.get("error")] if raw_items else []

        posts_for_group = []
        raw_count = len(items)
        group_posts_count[g["group_id"]] = raw_count

        if items:
            # Billing: charge AFTER successful scan that returned data
            if not billing.charge_scan(g["group_id"], g["nazwa"]):
                groups_skipped_billing += 1
                logger.warning("Group %s: skipped (billing)", g["nazwa"][:40])
                group_post_map[g["group_id"]] = []
                continue

            groups_ok += 1
            for item in items:
                post = map_item(item, g["nazwa"])
                if post:
                    all_posts.append(post)
                    posts_for_group.append(post)
        else:
            groups_failed += 1
            logger.warning("Group %s: no data (not charged)", g["nazwa"][:40])
            gm.record_scan(g["group_id"], queries=0, hits=0, posts_returned=0)
        group_post_map[g["group_id"]] = posts_for_group

    # End-of-scraping guard check (auth mode only)
    if use_auth:
        guard.end_of_cycle_check()

    total_fetched = len(all_posts)

    # Freshness filter
    cutoff = datetime.now() - timedelta(hours=settings.post_max_age_hours)
    fresh = [p for p in all_posts if p.timestamp.replace(tzinfo=None) >= cutoff]
    skipped_old = total_fetched - len(fresh)
    all_posts = fresh

    logger.info("Posts: fetched=%d too_old=%d to_analyze=%d", total_fetched, skipped_old, len(all_posts))

    # Pipeline
    index = CatalogIndex(settings.db_path)
    dgred = DgredClient(settings)
    writer = LeadWriter(dgred, settings)

    statuses: Counter[str] = Counter()
    errors = 0
    leads_written = 0
    queries_found = 0
    duplicates = 0
    halted_reason = None

    # Track per-group stats for record_scan
    group_queries: Counter[str] = Counter()  # group_name → query count
    group_hits: Counter[str] = Counter()     # group_name → hit count

    # Load seen IDs for dedup
    seen_path = Path(settings.db_path).parent / "seen_post_ids.json"
    try:
        seen_ids = set(json.loads(seen_path.read_text())) if seen_path.exists() else set()
    except Exception:
        seen_ids = set()

    for post in all_posts:
        # Dedup
        if post.post_id in seen_ids:
            duplicates += 1
            continue

        # Extract
        try:
            query = extract(post, settings)
        except Exception:
            errors += 1
            logger.exception("Extraction error for %s", post.post_id)
            continue

        seen_ids.add(post.post_id)

        if query is None or not query.is_part_request:
            continue

        queries_found += 1

        # Find which group this post belongs to
        post_group_name = post.group

        # Match
        try:
            result = match(query, index)
        except Exception:
            errors += 1
            logger.exception("Match error for %s", post.post_id)
            continue

        status_name = map_status(result, settings)
        statuses[status_name] += 1

        # Track per-group stats
        group_queries[post_group_name] += 1
        if status_name in ("Dopasowano", "Do weryfikacji"):
            group_hits[post_group_name] += 1

        # Circuit breaker: check before writing
        if leads_written >= settings.max_leads_per_cycle:
            halted_reason = f"max_leads_per_cycle ({settings.max_leads_per_cycle}) reached"
            logger.warning("CIRCUIT BREAKER: %s", halted_reason)
            break

        if errors >= settings.max_errors_per_cycle:
            halted_reason = f"max_errors_per_cycle ({settings.max_errors_per_cycle}) reached"
            logger.warning("CIRCUIT BREAKER: %s", halted_reason)
            break

        # Billing: charge BEFORE writing lead
        if not billing.charge_lead(status_name, "pending", post.post_id):
            logger.warning("Lead for %s skipped (billing)", post.post_id)
            continue

        # Check billing circuit breaker (includes token exhaustion)
        if billing.halted:
            if billing.is_tokens_exhausted():
                halted_reason = "tokens exhausted"
                billing.send_tokens_alarm(settings, exhausted=True, balance=0)
                logger.error("TOKENS EXHAUSTED mid-cycle — halting")
            else:
                halted_reason = "billing circuit breaker"
                logger.warning("BILLING CIRCUIT BREAKER — halting cycle")
            break

        # Write to dgred
        try:
            writer.write_lead(post, query, result)
            leads_written += 1
        except Exception:
            errors += 1
            logger.exception("dgred write error for %s", post.post_id)

            if errors >= settings.max_errors_per_cycle:
                halted_reason = f"max_errors_per_cycle ({settings.max_errors_per_cycle}) — dgred errors"
                logger.warning("CIRCUIT BREAKER: %s", halted_reason)
                break

    # Persist seen IDs
    seen_path.parent.mkdir(parents=True, exist_ok=True)
    seen_path.write_text(json.dumps(sorted(seen_ids)))

    # Record per-group scan stats
    for g in groups:
        gname = g["nazwa"]
        if group_post_map.get(g["group_id"]):  # only if we got posts
            gm.record_scan(
                g["group_id"],
                queries=group_queries.get(gname, 0),
                hits=group_hits.get(gname, 0),
                posts_returned=group_posts_count.get(g["group_id"], 0),
            )

    # Auto-cull: recompute group states (untested → proven/dead)
    cull_result = gm.auto_cull()
    if cull_result.get("culled_this_run"):
        logger.info("Auto-cull: %d group(s) marked dead", cull_result["culled_this_run"])
    logger.info(
        "Funnel: proven=%d untested=%d dead=%d",
        cull_result.get("proven", 0),
        cull_result.get("untested", 0),
        cull_result.get("dead", 0),
    )

    # Update the cycle_log row reserved at start
    conn_upd = sqlite3.connect(settings.db_path)
    conn_upd.execute("""
        UPDATE cycle_log SET
            finished_at=?, groups_scanned=?, groups_failed=?,
            posts_fetched=?, posts_too_old=?, posts_analyzed=?,
            queries_found=?, leads_written=?, leads_matched=?,
            leads_review=?, leads_no_match=?, duplicates=?,
            errors=?, halted_reason=?
        WHERE id=?
    """, (
        datetime.now().isoformat(), groups_ok, groups_failed,
        total_fetched, skipped_old, len(all_posts),
        queries_found, leads_written,
        statuses.get("Dopasowano", 0), statuses.get("Do weryfikacji", 0),
        statuses.get("Brak dopasowania", 0), duplicates,
        errors, halted_reason, cycle_id,
    ))
    conn_upd.commit()
    conn_upd.close()

    billing.close()
    guard.close()
    index.close()
    dgred.close()
    gm.close()

    logger.info(
        "=== Cycle done: groups=%d/%d posts=%d queries=%d leads=%d "
        "(67=%d 68=%d 69=%d) errors=%d halted=%s ===",
        groups_ok, groups_ok + groups_failed, total_fetched,
        queries_found, leads_written,
        statuses.get("Dopasowano", 0), statuses.get("Do weryfikacji", 0),
        statuses.get("Brak dopasowania", 0),
        errors, halted_reason or "no",
    )

    return 1 if halted_reason else 0


if __name__ == "__main__":
    sys.exit(main())
