"""
Pick n Pay (pnp.co.za) product scraper.

Pipeline:
  1. fetch-sitemaps : download 7 product sitemaps, populate seeds table
  2. scrape        : iterate pending seeds, hit Hybris OCC API, store product JSON
  3. export        : dump finished products to JSONL / CSV
  4. status        : print queue stats

The site is SAP Commerce Cloud (Hybris) + Spartacus. The product API is:
  https://www.pnp.co.za/pnphybris/v2/pnp-spa/products/{code}?fields=FULL&storeCode={STORE}

Resumable: state lives in SQLite. Crash, kill, rerun — picks up where it left off.
"""

from __future__ import annotations

import argparse
import json
import os
import random
import re
import sqlite3
import sys
import time
import xml.etree.ElementTree as ET
from contextlib import closing
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Iterator

import requests

import config
from db_init import PNP_SCHEMA as SCHEMA, init_pnp_db

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

ROOT = config.ROOT
DB_PATH = config.PNP_DB
SITEMAP_INDEX = "https://www.pnp.co.za/sitemap.xml"
API_BASE = "https://www.pnp.co.za/pnphybris/v2/pnp-spa"
DEFAULT_STORE_CODE = "WC27"  # Pick n Pay Waterfront (flagship Cape Town)
DEFAULT_FIELDS = (
    "DEFAULT,averageRating,images(FULL),classifications,manufacturer,"
    "numberOfReviews,categories(FULL),baseOptions,baseProduct,"
    "variantOptions,variantType,quantityType"
)
SITEMAP_NS = {"sm": "http://www.sitemaps.org/schemas/sitemap/0.9",
              "image": "http://www.google.com/schemas/sitemap-image/1.1"}
USER_AGENT = ("Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 "
              "(KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36")
DEFAULT_DELAY = config.PNP_DELAY
DEFAULT_JITTER = 0.4         # +/- this fraction of the delay
MAX_ATTEMPTS = 4             # per product, with exponential backoff
# Gateway / rate-limit responses — retry with backoff (502 = PnP upstream overload)
RETRYABLE_HTTP = frozenset({429, 502, 503, 504})

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

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


def init_db() -> Path:
    """Create pnp.db with schema if missing."""
    return init_pnp_db()


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


# ------------------------------------------------------------------ http session

def make_session() -> requests.Session:
    s = requests.Session()
    s.headers.update({
        "User-Agent": USER_AGENT,
        "Accept": "application/json",
        "Accept-Language": "en-ZA,en;q=0.9",
        "Referer": "https://www.pnp.co.za/",
    })
    # warm up — site sets a `route` cookie that downstream requests need
    s.get("https://www.pnp.co.za/", timeout=30,
          headers={"Accept": "text/html,application/xhtml+xml"})
    return s


# ------------------------------------------------------------------ sitemaps

PRODUCT_URL_RE = re.compile(r"/([^/]+)/p/([A-Za-z0-9_]+)$")


def cmd_fetch_sitemaps(args: argparse.Namespace) -> None:
    """Download all PRODUCT sub-sitemaps and insert pending seeds."""
    session = requests.Session()
    session.headers["User-Agent"] = USER_AGENT

    print(f"[sitemap] fetching index: {SITEMAP_INDEX}")
    idx = session.get(SITEMAP_INDEX, timeout=30)
    idx.raise_for_status()
    root = ET.fromstring(idx.text)

    product_sitemaps = [
        loc.text for loc in root.findall(".//sm:sitemap/sm:loc", SITEMAP_NS)
        if "PRODUCT-" in (loc.text or "")
    ]
    print(f"[sitemap] found {len(product_sitemaps)} product sub-sitemaps")

    conn = db_connect()
    total_new = 0
    total_seen = 0
    with conn:
        for i, sm_url in enumerate(product_sitemaps):
            print(f"[sitemap] {i+1}/{len(product_sitemaps)} GET {sm_url}")
            r = session.get(sm_url, timeout=120)
            r.raise_for_status()
            sub = ET.fromstring(r.content)
            batch = []
            for url_el in sub.findall("sm:url", SITEMAP_NS):
                loc = (url_el.find("sm:loc", SITEMAP_NS).text or "").strip()
                m = PRODUCT_URL_RE.search(loc)
                if not m:
                    continue
                slug, code = m.group(1), m.group(2)
                img_el = url_el.find("image:image/image:loc", SITEMAP_NS)
                image_url = img_el.text.strip() if img_el is not None and img_el.text else None
                batch.append((code, loc, slug, image_url))
            total_seen += len(batch)
            cur = conn.executemany(
                "INSERT OR IGNORE INTO seeds(code,url,slug,sitemap_image_url) "
                "VALUES (?,?,?,?)",
                batch,
            )
            total_new += cur.rowcount
            print(f"[sitemap]  parsed {len(batch):>6} urls, new seeds: {cur.rowcount}")

    pending = conn.execute("SELECT COUNT(*) FROM seeds WHERE status='pending'").fetchone()[0]
    total = conn.execute("SELECT COUNT(*) FROM seeds").fetchone()[0]
    print(f"[sitemap] done. urls seen: {total_seen}, new inserts: {total_new}")
    print(f"[sitemap] seeds total: {total}, pending: {pending}")


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

def category_path(categories: list[dict]) -> str:
    return " > ".join((c.get("name") or c.get("code") or "?") for c in categories or [])


def primary_image(images: list[dict]) -> str | None:
    if not images:
        return None
    # prefer largest PRIMARY zoom format
    for img in images:
        if img.get("imageType") == "PRIMARY" and img.get("format") == "zoom":
            return img.get("url")
    for img in images:
        if img.get("imageType") == "PRIMARY":
            return img.get("url")
    return images[0].get("url")


def _average_weight_value(raw) -> float | None:
    """Hybris returns averageWeight as a number or {value, formatedAverageWeight}."""
    if raw is None:
        return None
    if isinstance(raw, dict):
        v = raw.get("value")
        return float(v) if v is not None else None
    if isinstance(raw, (int, float)):
        return float(raw)
    return None


def extract_barcode(j: dict) -> str | None:
    """Pull unit_barcode out of the Hybris classifications block."""
    for cls in j.get("classifications") or []:
        for f in cls.get("features") or []:
            if (f.get("name") or "").lower() == "unit_barcode":
                fv = f.get("featureValues") or []
                if fv and fv[0].get("value"):
                    return str(fv[0]["value"]).strip()
    return None


def extract_summary_row(code: str, j: dict) -> tuple:
    price = j.get("price") or {}
    stock = j.get("stock") or {}
    cats = j.get("categories") or []
    imgs = j.get("images") or []
    mfr = j.get("manufacturer")
    if isinstance(mfr, dict):
        mfr = mfr.get("name") or mfr.get("code")
    return (
        code,
        j.get("name"),
        j.get("brand"),
        extract_barcode(j),
        mfr,
        price.get("value"),
        price.get("currencyIso"),
        price.get("formattedValue"),
        j.get("defaultUnitOfMeasure"),
        j.get("quantityType"),
        _average_weight_value(j.get("averageWeight")),
        j.get("inStockIndicator"),
        stock.get("stockLevelStatus"),
        stock.get("stockLevel"),
        1 if j.get("available") else 0,
        1 if j.get("purchasable") else 0,
        j.get("averageRating"),
        j.get("numberOfReviews"),
        j.get("summary"),
        j.get("description"),
        primary_image(imgs),
        len(imgs),
        category_path(cats),
        ",".join(c.get("code", "") for c in cats),
        j.get("url"),
        json.dumps(j, separators=(",", ":")),
        now_iso(),
    )


PRODUCT_COLS = (
    "code,name,brand,barcode,manufacturer,price_value,price_currency,price_formatted,"
    "unit_of_measure,quantity_type,average_weight,in_stock_indicator,"
    "stock_status,stock_level,available,purchasable,average_rating,"
    "number_of_reviews,summary,description,primary_image_url,image_count,"
    "category_path,category_codes,url,raw_json,scraped_at"
)
PRODUCT_QMARKS = ",".join(["?"] * len(PRODUCT_COLS.split(",")))


def iter_pending(conn: sqlite3.Connection, batch: int = 200,
                 retry_errors: bool = False) -> Iterator[tuple[str, str]]:
    statuses = "('pending','error')" if retry_errors else "('pending')"
    while True:
        rows = conn.execute(
            f"SELECT code, url FROM seeds "
            f"WHERE status IN {statuses} AND attempts < ? "
            f"ORDER BY attempts ASC, code ASC LIMIT ?",
            (MAX_ATTEMPTS, batch),
        ).fetchall()
        if not rows:
            return
        for r in rows:
            yield r[0], r[1]


def fetch_product(session: requests.Session, code: str, store_code: str,
                  fields: str, timeout: int = 30) -> tuple[int, dict | None, str | None]:
    url = f"{API_BASE}/products/{code}?fields={fields}&storeCode={store_code}"
    try:
        r = session.get(url, timeout=timeout)
    except requests.RequestException as e:
        return 0, None, f"network: {e!r}"
    if r.status_code == 200 and "json" in r.headers.get("Content-Type", ""):
        try:
            return r.status_code, r.json(), None
        except ValueError as e:
            return r.status_code, None, f"json-decode: {e!r}"
    if r.status_code == 404:
        return r.status_code, None, "not-found"
    return r.status_code, None, f"http-{r.status_code}: {r.text[:200]}"


def cmd_scrape(args: argparse.Namespace) -> None:
    conn = db_connect()
    if args.refresh:
        with conn:
            n = conn.execute(
                "UPDATE seeds SET status='pending' WHERE status='done'"
            ).rowcount
        print(f"[scrape] refresh: re-queued {n} done seeds")
    pending = conn.execute("SELECT COUNT(*) FROM seeds WHERE status='pending'").fetchone()[0]
    if pending == 0 and not args.retry_errors:
        print("[scrape] nothing pending. (use --retry-errors to retry failures)")
        return
    print(f"[scrape] starting. pending={pending} store={args.store} delay={args.delay}s")

    session = make_session()
    delay = args.delay
    jitter = args.jitter
    fields = args.fields
    limit = args.limit
    done = 0
    err = 0
    notfound = 0
    gateway_streak = 0

    try:
        for code, _url in iter_pending(conn, retry_errors=args.retry_errors):
            if limit and done + err + notfound >= limit:
                print(f"[scrape] hit --limit {limit}")
                break

            # exponential backoff per-attempt
            attempt_row = conn.execute(
                "SELECT attempts FROM seeds WHERE code=?", (code,)
            ).fetchone()
            attempts = attempt_row[0] if attempt_row else 0
            if attempts > 0:
                sleep_for = min(60, delay * (2 ** attempts)) + random.uniform(0, 1)
                time.sleep(sleep_for)
            else:
                time.sleep(delay * (1 + random.uniform(-jitter, jitter)))

            status, payload, err_msg = fetch_product(session, code, args.store, fields)

            if payload is not None:
                row = extract_summary_row(code, payload)
                with conn:
                    conn.execute(
                        f"INSERT OR REPLACE INTO products({PRODUCT_COLS}) "
                        f"VALUES ({PRODUCT_QMARKS})",
                        row,
                    )
                    conn.execute(
                        "UPDATE seeds SET status='done', attempts=attempts+1, "
                        "http_status=?, fetched_at=?, last_error=NULL WHERE code=?",
                        (status, now_iso(), code),
                    )
                done += 1
                if done % 50 == 0:
                    print(f"[scrape] done={done} err={err} 404={notfound}  "
                          f"last: {payload.get('name','?')[:60]}")
            elif status == 404:
                with conn:
                    conn.execute(
                        "UPDATE seeds SET status='notfound', attempts=attempts+1, "
                        "http_status=404, fetched_at=?, last_error=? WHERE code=?",
                        (now_iso(), err_msg, code),
                    )
                notfound += 1
            else:
                # transient — leave as 'pending' if attempts < MAX, else 'error'
                new_status = "error" if (attempts + 1) >= MAX_ATTEMPTS else "pending"
                with conn:
                    conn.execute(
                        "UPDATE seeds SET status=?, attempts=attempts+1, "
                        "http_status=?, fetched_at=?, last_error=? WHERE code=?",
                        (new_status, status, now_iso(), err_msg, code),
                    )
                err += 1
                if err % 10 == 0:
                    print(f"[scrape]  recent failure ({status}): {err_msg[:120] if err_msg else ''}")
                if status in RETRYABLE_HTTP:
                    gateway_streak += 1
                    cool = 30 + random.uniform(0, 30) * min(gateway_streak, 4)
                    print(f"[scrape] {status} received — cooling down {cool:.0f}s")
                    time.sleep(cool)
                    if gateway_streak >= 5:
                        print("[scrape] refreshing session after repeated gateway errors")
                        session = make_session()
                        gateway_streak = 0
                else:
                    gateway_streak = 0
    except KeyboardInterrupt:
        print("\n[scrape] interrupted by user — progress saved.")

    print(f"[scrape] finished. done={done} err={err} 404={notfound}")


# ------------------------------------------------------------------ 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")

    # full JSONL
    jsonl_path = out_dir / f"products_full_{stamp}.jsonl"
    with closing(conn.execute("SELECT raw_json FROM products")) as cur, \
            open(jsonl_path, "w", encoding="utf-8") as f:
        n = 0
        for (raw,) in cur:
            f.write(raw)
            f.write("\n")
            n += 1
    print(f"[export] wrote {n} products to {jsonl_path}")

    # flat summary JSONL
    flat_path = out_dir / f"products_flat_{stamp}.jsonl"
    cols = [c for c in PRODUCT_COLS.split(",") if c != "raw_json"]
    sql = f"SELECT {','.join(cols)} FROM products"
    with closing(conn.execute(sql)) as cur, \
            open(flat_path, "w", encoding="utf-8") as f:
        n = 0
        for row in cur:
            f.write(json.dumps(dict(zip(cols, row)), ensure_ascii=False))
            f.write("\n")
            n += 1
    print(f"[export] wrote {n} flat rows to {flat_path}")

    # CSV
    if args.csv:
        import csv
        csv_path = out_dir / f"products_flat_{stamp}.csv"
        with closing(conn.execute(sql)) as cur, \
                open(csv_path, "w", encoding="utf-8", newline="") as f:
            w = csv.writer(f)
            w.writerow(cols)
            w.writerows(cur)
        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 fetch-sitemaps first.")
        return
    conn = db_connect()
    print(f"DB: {DB_PATH} ({DB_PATH.stat().st_size/1024/1024:.1f} MB)")
    print("--- seeds by status ---")
    for status, count in conn.execute(
            "SELECT status, COUNT(*) FROM seeds GROUP BY status ORDER BY 2 DESC"):
        print(f"  {status:<10} {count:>8}")
    total = conn.execute("SELECT COUNT(*) FROM seeds").fetchone()[0]
    products = conn.execute("SELECT COUNT(*) FROM products").fetchone()[0]
    print(f"--- totals ---")
    print(f"  seeds:    {total}")
    print(f"  products: {products}")
    # recent errors
    print("--- recent errors (last 5) ---")
    for code, status, http, err in conn.execute(
            "SELECT code, status, http_status, last_error FROM seeds "
            "WHERE status IN ('error','notfound') ORDER BY fetched_at DESC LIMIT 5"):
        print(f"  {code} [{status}/{http}] {(err or '')[:100]}")


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

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

    sp = sub.add_parser("fetch-sitemaps", help="download sitemaps and seed the queue")
    sp.set_defaults(func=cmd_fetch_sitemaps)

    sp = sub.add_parser("scrape", help="scrape pending products via the OCC API")
    sp.add_argument("--store", default=DEFAULT_STORE_CODE,
                    help=f"Hybris storeCode (default {DEFAULT_STORE_CODE} = Waterfront)")
    sp.add_argument("--delay", type=float, default=DEFAULT_DELAY,
                    help="seconds between requests (default 1.0)")
    sp.add_argument("--jitter", type=float, default=DEFAULT_JITTER,
                    help="random jitter as fraction of delay (default 0.4)")
    sp.add_argument("--fields", default=DEFAULT_FIELDS, help="OCC fields parameter")
    sp.add_argument("--limit", type=int, default=0,
                    help="stop after N products (0 = unlimited)")
    sp.add_argument("--retry-errors", action="store_true",
                    help="also retry rows currently in 'error' state")
    sp.add_argument("--refresh", action="store_true",
                    help="re-queue done seeds so prices are re-fetched")
    sp.set_defaults(func=cmd_scrape)

    sp = sub.add_parser("export", help="dump products to JSONL/CSV")
    sp.add_argument("--csv", action="store_true", help="also write flat CSV")
    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()
