"""
Checkers (checkers.co.za / Sixty60) product scraper — hybrid approach.

The Checkers storefront is a Next.js SPA behind AWS WAF. Its product API is a
same-origin POST the SPA makes:

  POST https://www.checkers.co.za/api/catalogue/get-products-filter
  body: {... "productListSource": {"search": "<term>"} ...}

It returns rich JSON per product including `barcodes` (real EAN/GTIN — our
cross-retailer join key), `price`, `oldPrice`, `name`, stock, images, etc.

To get past the AWS WAF we drive a real Chromium via Playwright to obtain the
`aws-waf-token` cookie, then replay the JSON API through Playwright's request
context (which carries the cookie and re-solves the challenge if it expires).
No page rendering per product -> fast.

Seeding: the shared MVP staples list in mvp_terms.py (real household shopping
behaviour, ~1,000 SKUs after dedup), NOT a full-store crawl.

Resumable: state lives in checkers.db (SQLite). Crash/kill/rerun continues.

Commands:
  scrape  : iterate pending search terms, collect products
  export  : dump products to JSONL / CSV
  status  : queue stats
"""
from __future__ import annotations

import argparse
import json
import random
import sqlite3
import sys
import time
from datetime import datetime, timezone
from pathlib import Path

from playwright.sync_api import sync_playwright

import config
import mvp_terms

# ------------------------------------------------------------------ config

ROOT = config.ROOT
DB_PATH = config.CHECKERS_DB
HOME = "https://www.checkers.co.za/"
API = "https://www.checkers.co.za/api/catalogue/get-products-filter"
UA = ("Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 "
      "(KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36")
PAGE_SIZE = 50
PER_TERM_CAP = config.PER_TERM_CAP
DEFAULT_DELAY = config.CHECKERS_DELAY
RETAILER = "checkers"

# ------------------------------------------------------------------ db

SCHEMA = """
CREATE TABLE IF NOT EXISTS terms (
    bucket      TEXT NOT NULL,
    term        TEXT NOT NULL,
    status      TEXT NOT NULL DEFAULT 'pending',
    total_count INTEGER,
    n_collected INTEGER DEFAULT 0,
    attempts    INTEGER NOT NULL DEFAULT 0,
    last_error  TEXT,
    fetched_at  TEXT,
    PRIMARY KEY (bucket, term)
);

CREATE TABLE IF NOT EXISTS products (
    id                 TEXT PRIMARY KEY,
    barcode            TEXT,
    name               TEXT,
    brand              TEXT,
    price_value        REAL,
    was_price          REAL,
    on_promotion       INTEGER,
    currency           TEXT,
    unit_of_measure    TEXT,
    pack_quantity      INTEGER,
    stock_on_hand      INTEGER,
    out_of_stock       INTEGER,
    article_number     TEXT,
    merchandise_cat    TEXT,
    primary_image_url  TEXT,
    store_id           TEXT,
    buckets            TEXT,
    matched_terms      TEXT,
    url                TEXT,
    raw_json           TEXT NOT NULL,
    scraped_at         TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_ck_barcode ON products(barcode);
CREATE INDEX IF NOT EXISTS idx_ck_name ON products(name);
"""

PRODUCT_COLS = (
    "id,barcode,name,brand,price_value,was_price,on_promotion,currency,"
    "unit_of_measure,pack_quantity,stock_on_hand,out_of_stock,article_number,"
    "merchandise_cat,primary_image_url,store_id,buckets,matched_terms,url,"
    "raw_json,scraped_at"
)
PRODUCT_QMARKS = ",".join(["?"] * len(PRODUCT_COLS.split(",")))


def db_connect() -> sqlite3.Connection:
    conn = sqlite3.connect(DB_PATH)
    conn.execute("PRAGMA journal_mode=WAL")
    conn.execute("PRAGMA synchronous=NORMAL")
    conn.executescript(SCHEMA)
    return conn


def now_iso() -> str:
    return datetime.now(timezone.utc).isoformat(timespec="seconds")


def seed_terms(conn: sqlite3.Connection) -> None:
    with conn:
        conn.executemany(
            "INSERT OR IGNORE INTO terms(bucket, term) VALUES (?, ?)",
            mvp_terms.all_terms(),
        )


# ------------------------------------------------------------------ extraction

def search_payload(term: str, page: int, page_size: int) -> dict:
    return {
        "storeContexts": [],
        "filterData": {
            "filter": {
                "showAllDisplayVariants": False,
                "showNotRangedProducts": False,
                "productListSource": {"search": term},
                "paginationOptions": {"page": page, "pageSize": page_size},
                "filterOptions": {"filterIds": [], "dealsOnly": False,
                                  "brandOptions": [], "departmentOptions": [],
                                  "serviceOptions": [], "facetOptions": []},
                "sortOptions": None,
            },
            "displayOptions": {"includeDisplayCategoryTree": False},
        },
        "forYouBonusBuyIds": [],
    }


def extract_row(p: dict, bucket: str, term: str) -> tuple:
    barcodes = p.get("barcodes") or []
    barcode = str(barcodes[0]).strip() if barcodes else None
    factor = p.get("priceFactor") or 100
    old = p.get("oldPrice")
    was = (old / factor) if isinstance(old, (int, float)) and old else None
    price = p.get("price")
    if price is None and isinstance(p.get("priceWithoutDecimal"), (int, float)):
        price = p["priceWithoutDecimal"] / factor
    article = p.get("articleNumber")
    slug_uom = (p.get("unitOfMeasure") or "")
    url = (f"https://www.checkers.co.za/product/{article}{slug_uom}"
           if article else None)
    return (
        p.get("id"),
        barcode,
        p.get("name") or p.get("displayName"),
        None,  # brand: not a discrete field; derived later in matching from name
        price,
        was if (was and price is not None and was > price) else None,
        1 if p.get("isOnPromotion") else 0,
        p.get("currency"),
        p.get("unitOfMeasure"),
        p.get("packQuantity"),
        p.get("stockOnHand"),
        1 if p.get("outOfStock") else 0,
        str(article) if article is not None else None,
        str(p.get("merchandiseCategory")) if p.get("merchandiseCategory") else None,
        p.get("imageProductCardURL") or p.get("imageURL"),
        p.get("storeId"),
        bucket,
        term,
        url,
        json.dumps(p, separators=(",", ":")),
        now_iso(),
    )


def upsert_products(conn: sqlite3.Connection, products: list[dict],
                    bucket: str, term: str) -> int:
    """Insert new products; for existing ids, append bucket/term provenance."""
    new = 0
    with conn:
        for p in products:
            pid = p.get("id")
            if not pid:
                continue
            existing = conn.execute(
                "SELECT buckets, matched_terms FROM products WHERE id=?", (pid,)
            ).fetchone()
            if existing:
                buckets = set(filter(None, (existing[0] or "").split(",")))
                terms = set(filter(None, (existing[1] or "").split("|")))
                buckets.add(bucket)
                terms.add(term)
                conn.execute(
                    "UPDATE products SET buckets=?, matched_terms=? WHERE id=?",
                    (",".join(sorted(buckets)), "|".join(sorted(terms)), pid),
                )
            else:
                conn.execute(
                    f"INSERT INTO products({PRODUCT_COLS}) VALUES ({PRODUCT_QMARKS})",
                    extract_row(p, bucket, term),
                )
                new += 1
    return new


# ------------------------------------------------------------------ scrape

class Session:
    """Playwright browser context that passes WAF and replays the JSON API."""

    def __init__(self, pw, headless=True):
        self.br = pw.chromium.launch(headless=headless)
        self.ctx = self.br.new_context(user_agent=UA, locale="en-ZA")
        self.page = self.ctx.new_page()
        self._warm()

    def _warm(self):
        self.page.goto(HOME, wait_until="networkidle", timeout=60000)
        self.page.wait_for_timeout(2000)

    def search(self, term: str, page: int) -> tuple[list[dict], int]:
        resp = self.ctx.request.post(
            API,
            data=json.dumps(search_payload(term, page, PAGE_SIZE)),
            headers={"content-type": "application/json", "accept": "application/json",
                     "origin": "https://www.checkers.co.za", "referer": HOME},
            timeout=45000,
        )
        if resp.status != 200:
            # likely WAF expired -> re-warm once and retry
            if resp.status in (403, 405, 401):
                self._warm()
                resp = self.ctx.request.post(
                    API, data=json.dumps(search_payload(term, page, PAGE_SIZE)),
                    headers={"content-type": "application/json", "accept": "application/json",
                             "origin": "https://www.checkers.co.za", "referer": HOME},
                    timeout=45000)
            if resp.status != 200:
                raise RuntimeError(f"http {resp.status}: {resp.text()[:160]}")
        j = resp.json()
        return j.get("products") or [], int(j.get("totalCount") or 0)

    def close(self):
        self.br.close()


def cmd_scrape(args: argparse.Namespace) -> None:
    conn = db_connect()
    seed_terms(conn)
    if args.refresh:
        with conn:
            n = conn.execute(
                "UPDATE terms SET status='pending' WHERE status='done'"
            ).rowcount
        print(f"[scrape] refresh: re-queued {n} done terms")
    statuses = "('pending','error')" if args.retry_errors else "('pending')"
    pending = conn.execute(
        f"SELECT bucket, term FROM terms WHERE status IN {statuses} "
        f"ORDER BY status DESC, bucket, term"
    ).fetchall()
    if not pending:
        print("[scrape] nothing pending. (use --retry-errors to retry failures)")
        return
    if args.limit:
        pending = pending[: args.limit]
    print(f"[scrape] {len(pending)} terms pending  delay={args.delay}s  "
          f"page_size={PAGE_SIZE}  cap={PER_TERM_CAP}")

    total_new = 0
    with sync_playwright() as pw:
        sess = Session(pw, headless=not args.headed)
        try:
            for i, (bucket, term) in enumerate(pending, 1):
                try:
                    collected = 0
                    new_here = 0
                    page = 0
                    total = None
                    while collected < PER_TERM_CAP:
                        products, total = sess.search(term, page)
                        if not products:
                            break
                        new_here += upsert_products(conn, products, bucket, term)
                        collected += len(products)
                        page += 1
                        if collected >= total or len(products) < PAGE_SIZE:
                            break
                        time.sleep(args.delay * (1 + random.uniform(-0.3, 0.3)))
                    with conn:
                        conn.execute(
                            "UPDATE terms SET status='done', total_count=?, "
                            "n_collected=?, attempts=attempts+1, fetched_at=?, "
                            "last_error=NULL WHERE bucket=? AND term=?",
                            (total, collected, now_iso(), bucket, term))
                    total_new += new_here
                    print(f"[scrape] {i}/{len(pending)} {bucket}/{term!r}: "
                          f"+{new_here} new (saw {collected}/{total}), total_new={total_new}")
                except Exception as e:
                    with conn:
                        conn.execute(
                            "UPDATE terms SET status='error', attempts=attempts+1, "
                            "fetched_at=?, last_error=? WHERE bucket=? AND term=?",
                            (now_iso(), repr(e)[:300], bucket, term))
                    print(f"[scrape] {i}/{len(pending)} {bucket}/{term!r}: ERROR {e}")
                time.sleep(args.delay * (1 + random.uniform(-0.3, 0.3)))
        except KeyboardInterrupt:
            print("\n[scrape] interrupted — progress saved.")
        finally:
            sess.close()

    n = conn.execute("SELECT COUNT(*) FROM products").fetchone()[0]
    bc = conn.execute("SELECT COUNT(*) FROM products WHERE barcode IS NOT NULL").fetchone()[0]
    print(f"[scrape] done. products={n} with_barcode={bc} new_this_run={total_new}")


# ------------------------------------------------------------------ export

def cmd_export(args: argparse.Namespace) -> None:
    conn = db_connect()
    out_dir = config.EXPORTS_DIR
    out_dir.mkdir(exist_ok=True)
    stamp = datetime.now().strftime("%Y%m%d_%H%M%S")
    cols = [c for c in PRODUCT_COLS.split(",") if c != "raw_json"]
    sql = f"SELECT {','.join(cols)} FROM products"

    jsonl_path = out_dir / f"checkers_flat_{stamp}.jsonl"
    n = 0
    with open(jsonl_path, "w", encoding="utf-8") as f:
        for row in conn.execute(sql):
            f.write(json.dumps(dict(zip(cols, row)), ensure_ascii=False) + "\n")
            n += 1
    print(f"[export] wrote {n} rows to {jsonl_path}")

    if args.csv:
        import csv
        csv_path = out_dir / f"checkers_flat_{stamp}.csv"
        with open(csv_path, "w", encoding="utf-8", newline="") as f:
            w = csv.writer(f)
            w.writerow(cols)
            w.writerows(conn.execute(sql))
        print(f"[export] wrote CSV to {csv_path}")


# ------------------------------------------------------------------ status

def cmd_status(_args: argparse.Namespace) -> None:
    if not DB_PATH.exists():
        print(f"[status] no database at {DB_PATH}. Run scrape first.")
        return
    conn = db_connect()
    seed_terms(conn)
    print(f"DB: {DB_PATH} ({DB_PATH.stat().st_size/1024/1024:.1f} MB)")
    print("--- terms by status ---")
    for status, count in conn.execute(
            "SELECT status, COUNT(*) FROM terms GROUP BY status ORDER BY 2 DESC"):
        print(f"  {status:<10} {count:>6}")
    n = conn.execute("SELECT COUNT(*) FROM products").fetchone()[0]
    bc = conn.execute("SELECT COUNT(*) FROM products WHERE barcode IS NOT NULL").fetchone()[0]
    print(f"--- products: {n}  with_barcode: {bc} ---")
    print("--- by bucket ---")
    for bucket, count in conn.execute(
            "SELECT buckets, COUNT(*) FROM products GROUP BY buckets ORDER BY 2 DESC LIMIT 12"):
        print(f"  {bucket:<40} {count:>6}")


# ------------------------------------------------------------------ cli

def main() -> None:
    p = argparse.ArgumentParser(description="Checkers (Sixty60) product scraper")
    sub = p.add_subparsers(dest="cmd", required=True)

    sp = sub.add_parser("scrape", help="scrape MVP search terms")
    sp.add_argument("--delay", type=float, default=DEFAULT_DELAY)
    sp.add_argument("--limit", type=int, default=0, help="only first N terms (testing)")
    sp.add_argument("--retry-errors", action="store_true")
    sp.add_argument("--refresh", action="store_true",
                    help="re-queue done terms so prices are re-fetched")
    sp.add_argument("--headed", action="store_true", help="show the browser window")
    sp.set_defaults(func=cmd_scrape)

    sp = sub.add_parser("export", help="dump products to JSONL/CSV")
    sp.add_argument("--csv", action="store_true")
    sp.set_defaults(func=cmd_export)

    sp = sub.add_parser("status", help="show queue stats")
    sp.set_defaults(func=cmd_status)

    args = p.parse_args()
    args.func(args)


if __name__ == "__main__":
    main()
