"""Decides which products may be compared with which.

218 products compare in 47,000 pairs and would run fine; 20,000 compare in 400
million and would not. Blocking keeps the job's cost tied to catalog shape rather
than catalog size squared, without making cross-category complements impossible.
"""
import logging

from psycopg2 import sql
from psycopg2.extras import RealDictCursor

from app.services.catalog.normalise import BLOCKING_FIELDS
from app.services.infra.database import get_db_connection

logger = logging.getLogger(__name__)

UNCATEGORISED = "__uncategorised__"
MAX_CROSS_CATEGORY_PARTNERS = 40


def is_eligible(product: dict) -> bool:
    """Only a blocking gap disqualifies a product from being paired.

    Requiring an entirely empty missing_fields also excluded anything without a
    converted reference price, which quietly made currency mandatory on sources
    where it is optional -- a whole catalogue would ingest and pair nothing.
    """
    blocking = set(product.get("missing_fields") or []) & set(BLOCKING_FIELDS)
    return bool(
        product.get("in_stock")
        and not blocking
        and product.get("attributes")
    )


def blocks(products: list) -> dict:
    grouped = {}
    for product in products:
        grouped.setdefault(product.get("category") or UNCATEGORISED, []).append(product)
    return grouped


def candidate_pairs(products: list) -> list:
    """Every ordered pair worth scoring.

    Within a block: everything, both directions -- similar and upsell only ever
    apply here. Across blocks, two rules produce pairs: an accessory partner
    (a phone and a case), and two colour-carrying non-accessories (a shirt and
    trousers -- the outfit case, which has no accessory to key on). Comparing
    every phone against every grocery item buys nothing and is skipped.
    """
    grouped = blocks(products)
    pairs = []

    for members in grouped.values():
        for a in members:
            for b in members:
                if a["product_key"] != b["product_key"]:
                    pairs.append((a, b))

    accessories = [p for p in products if p.get("is_accessory")]
    non_accessories = [p for p in products if not p.get("is_accessory")]
    coloured = [p for p in non_accessories if _has_colour(p)]

    for anchor in non_accessories:
        anchor_category = anchor.get("category") or UNCATEGORISED
        # Truncated while collecting rather than after: building the full list
        # first would restore the O(n^2) pass this blocking exists to avoid, on
        # exactly the large catalogs where the cap matters most.
        partners = _take(
            (p for p in accessories
             if (p.get("category") or UNCATEGORISED) != anchor_category),
            MAX_CROSS_CATEGORY_PARTNERS)

        if _has_colour(anchor) and len(partners) < MAX_CROSS_CATEGORY_PARTNERS:
            partners += _take(
                (p for p in coloured
                 if p["product_key"] != anchor["product_key"]
                 and (p.get("category") or UNCATEGORISED) != anchor_category),
                MAX_CROSS_CATEGORY_PARTNERS - len(partners))

        pairs.extend((anchor, p) for p in partners)

    return pairs


def _take(candidates, limit: int) -> list:
    taken = []
    for candidate in candidates:
        if len(taken) >= limit:
            break
        taken.append(candidate)
    return taken


def _has_colour(product: dict) -> bool:
    return any(a.get("key") == "color" and a.get("value")
               for a in (product.get("attributes") or []))


def load_eligible(tenant_id: str) -> list:
    conn = get_db_connection()
    try:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            cur.execute(sql.SQL(
                "SELECT product_key, name, description, category, brand, "
                "attributes, price_cents, price_reference_cents, price_tier, "
                "is_accessory, rating, in_stock, missing_fields, content_hash, "
                "tenant_relations "
                "FROM {}.strategist_products"
            ).format(sql.Identifier(tenant_id)))
            rows = [dict(r) for r in cur.fetchall()]
    finally:
        conn.close()

    eligible = [r for r in rows if is_eligible(r)]
    if len(eligible) < len(rows):
        logger.info("%s: pairing %d of %d products; the rest are out of stock, "
                    "flagged, or have no attributes",
                    tenant_id, len(eligible), len(rows))
    return eligible
