"""Runs attribute extraction across a tenant's catalog.

Selection is driven by content_hash: a product is re-extracted only when its
content actually changed. That is the real cost control -- the matching
specification's quality-score gate was measured at a 0.35 median on live data,
so it would catch nearly everything and control nothing.
"""
import json
import logging

from psycopg2 import sql
from psycopg2.extras import Json, RealDictCursor

from app.services.enrichment.extractor import BATCH_SIZE, extract_batch
from app.services.enrichment.merge import merge_attributes
from app.services.infra.database import get_db_connection, migrate_products_table
from app.services.enrichment.price_tiers import classify, compute_bands

logger = logging.getLogger(__name__)

MIN_TEXT_WORDS = 3

# Crawled rows carry no content_hash at all. Storing NULL in enriched_hash for
# them would leave "extracted" indistinguishable from "never extracted", and
# every run would re-extract -- and re-bill -- the whole crawled set.
NO_CONTENT_HASH = "no-content-hash"


def is_extractable(product: dict) -> bool:
    """No prompt fixes an absent input, so do not pay for a call that can only
    guess."""
    if product.get("description"):
        return True
    name = product.get("name") or ""
    return len(name.split()) >= MIN_TEXT_WORDS


def select_products(tenant_id: str, force: bool) -> list:
    where = ("TRUE" if force else
             "enriched_hash IS NULL "
             "OR enriched_hash IS DISTINCT FROM COALESCE(content_hash, %s)")
    params = () if force else (NO_CONTENT_HASH,)
    conn = get_db_connection()
    try:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            cur.execute(sql.SQL(
                "SELECT product_key, name, description, brand, taxonomy_path, "
                "attributes, price_cents, price_reference_cents, content_hash "
                "FROM {}.strategist_products WHERE " + where
            ).format(sql.Identifier(tenant_id)), params)
            return [dict(r) for r in cur.fetchall()]
    finally:
        conn.close()


def catalog_stats(tenant_id: str) -> tuple:
    """(total products, products holding at least one attribute)."""
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(sql.SQL(
                "SELECT count(*), count(*) FILTER "
                "(WHERE jsonb_array_length(attributes) > 0) "
                "FROM {}.strategist_products"
            ).format(sql.Identifier(tenant_id)))
            return cur.fetchone()
    finally:
        conn.close()


def persist_attributes(tenant_id: str, product_key: str, attributes: list,
                       is_accessory, price_tier, content_hash) -> None:
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(sql.SQL(
                "UPDATE {}.strategist_products SET attributes = %s, "
                "is_accessory = COALESCE(%s, is_accessory), "
                "price_tier = %s, enriched_hash = %s WHERE product_key = %s"
            ).format(sql.Identifier(tenant_id)),
                (Json(attributes), is_accessory, price_tier,
                 content_hash or NO_CONTENT_HASH, product_key))
        conn.commit()
    except Exception:
        conn.rollback()
        logger.error("Could not persist attributes for %s", product_key,
                     exc_info=True)
        raise
    finally:
        conn.close()


def flag_unextractable(tenant_id: str, product_key: str) -> None:
    """Visible, not silent: a merchant should see a data gap rather than blame
    the recommendations."""
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(sql.SQL(
                "UPDATE {}.strategist_products "
                "SET missing_fields = array_append(missing_fields, 'attributes') "
                "WHERE product_key = %s AND NOT ('attributes' = ANY(missing_fields))"
            ).format(sql.Identifier(tenant_id)), (product_key,))
        conn.commit()
    except Exception:
        conn.rollback()
        logger.error("Could not flag %s as unextractable", product_key,
                     exc_info=True)
    finally:
        conn.close()


def _report(progress, percent: int, step: str) -> None:
    """Progress is best-effort. The job is spending real money on model calls;
    a reporting failure must never abort it."""
    if progress is None:
        return
    try:
        progress(int(percent), step)
    except Exception:
        logger.warning("Progress callback failed", exc_info=True)


def enrich_tenant(tenant_id: str, force: bool = False, progress=None) -> dict:
    migrate_products_table(tenant_id)
    products = select_products(tenant_id, force)
    bands = compute_bands(tenant_id)

    extractable = [p for p in products if is_extractable(p)]
    skipped = [p for p in products if not is_extractable(p)]
    for product in skipped:
        flag_unextractable(tenant_id, product["product_key"])

    extracted = failed = conflicts = 0
    total = len(extractable)

    if total == 0:
        _report(progress, 100, "extracting attributes — no products to extract")

    num_batches = -(-total // BATCH_SIZE)  # ceil division
    done = 0

    for start in range(0, total, BATCH_SIZE):
        batch = extractable[start:start + BATCH_SIZE]
        batch_number = start // BATCH_SIZE + 1
        results = extract_batch(batch)
        if not results:
            failed += len(batch)
            done += len(batch)
            _report(progress, done * 100 // total,
                    f"extracting attributes — batch {batch_number} of {num_batches}")
            continue

        for product in batch:
            inferred = results.get(product["product_key"])
            if not inferred:
                failed += 1
                continue

            current = product.get("attributes") or []
            if isinstance(current, str):
                current = json.loads(current)

            merged, product_conflicts = merge_attributes(current, inferred)
            conflicts += product_conflicts

            price = product.get("price_reference_cents") or product.get("price_cents")
            # The computed band wins over the model's guess: the model sees one
            # batch and cannot know the tenant's price distribution, which is
            # the only thing "premium" means here. Its guess is the fallback for
            # a product with no price at all.
            persist_attributes(
                tenant_id, product["product_key"], merged,
                inferred.get("is_accessory"),
                classify(price, bands) or inferred.get("price_tier"),
                product.get("content_hash"))
            extracted += 1

        done += len(batch)
        _report(progress, done * 100 // total,
                f"extracting attributes — batch {batch_number} of {num_batches}")

    total, with_attributes = catalog_stats(tenant_id)
    report = {"products": total,
              "considered": len(products),
              "extracted": extracted,
              "skipped_unchanged": total - len(products),
              "skipped_unextractable": len(skipped),
              "failed": failed,
              "conflicts": conflicts,
              "attribute_coverage": round(with_attributes / total, 3) if total else 0.0}
    logger.info("Enriched %s: %s", tenant_id, report)
    return report
