"""Read models behind the pairing endpoints.

The pairing job (pairing.py) and the decisions store (pairing_decisions.py) each
own their own table and their own writes; this module only reads, and joins the
two so the merchant-facing surface can show a pair together with its servability
without duplicating that rule.
"""
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
from app.services.pairing.decisions import (
    is_servable, load as load_decisions,
    load_for_anchor as load_decisions_for_anchor,
)
from app.services.pairing.rules import PAIR_TYPES

logger = logging.getLogger(__name__)

CARD_COLUMNS = ("name", "category", "price_cents", "currency", "image_url",
                "in_stock")


def _card(row: dict, prefix: str) -> dict:
    return {column: row.pop(f"{prefix}_{column}") for column in CARD_COLUMNS}


def list_categories(tenant_id: str) -> list:
    conn = get_db_connection()
    try:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            cur.execute(sql.SQL(
                "SELECT category, COUNT(*) AS products "
                "FROM {}.strategist_products "
                "GROUP BY category ORDER BY category"
            ).format(sql.Identifier(tenant_id)))
            return [dict(r) for r in cur.fetchall()]
    finally:
        conn.close()


def list_products(tenant_id: str, category: str = None, limit: int = 50,
                   offset: int = 0) -> dict:
    conn = get_db_connection()
    try:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            where = sql.SQL("WHERE category = %s") if category else sql.SQL("")
            params = [category] if category else []

            cur.execute(sql.SQL(
                "SELECT COUNT(*) AS total FROM {}.strategist_products {}"
            ).format(sql.Identifier(tenant_id), where), params)
            total = cur.fetchone()["total"]

            cur.execute(sql.SQL(
                "SELECT product_key, name, category, brand, price_cents, "
                "currency, in_stock, image_url FROM {}.strategist_products {} "
                "ORDER BY product_key LIMIT %s OFFSET %s"
            ).format(sql.Identifier(tenant_id), where), params + [limit, offset])
            products = [dict(r) for r in cur.fetchall()]
    finally:
        conn.close()
    return {"products": products, "total": total}


def get_product(tenant_id: str, product_key: str):
    conn = get_db_connection()
    try:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            cur.execute(sql.SQL(
                "SELECT * FROM {}.strategist_products WHERE product_key = %s"
            ).format(sql.Identifier(tenant_id)), (product_key,))
            row = cur.fetchone()
            return dict(row) if row else None
    finally:
        conn.close()


def pairings_for(tenant_id: str, product_key: str) -> dict:
    """Every neighbor of one anchor, grouped by type, each carrying `servable`.

    Servability is recomputed here rather than trusted from the job's report:
    a decision recorded after the last run must be reflected immediately, not
    only after the next re-pair.
    """
    conn = get_db_connection()
    try:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            cur.execute(sql.SQL("""
                SELECT n.anchor_key, n.neighbor_key, n.pair_type, n.score,
                       n.confidence, n.source, n.reasons, n.computed_at,
                       p.name AS neighbor_name, p.category AS neighbor_category,
                       p.price_cents AS neighbor_price_cents,
                       p.currency AS neighbor_currency,
                       p.image_url AS neighbor_image_url,
                       p.in_stock AS neighbor_in_stock
                FROM {}.strategist_product_neighbors n
                JOIN {}.strategist_products p ON p.product_key = n.neighbor_key
                WHERE n.anchor_key = %s
                ORDER BY n.score DESC
            """).format(sql.Identifier(tenant_id), sql.Identifier(tenant_id)),
                (product_key,))
            rows = [dict(r) for r in cur.fetchall()]
    finally:
        conn.close()

    decisions = load_decisions(tenant_id)
    grouped = {t: [] for t in PAIR_TYPES}
    for row in rows:
        row["neighbor"] = _card(row, "neighbor")
        decision = decisions.get(
            (row["anchor_key"], row["neighbor_key"], row["pair_type"]))
        row["servable"] = is_servable(row, decision)
        grouped.setdefault(row["pair_type"], []).append(row)
    return grouped


# Upsell is computed and visible in the admin tool but never served: it only
# makes sense at the moment a shopper is comparing, which needs the per-event
# targeting this system deliberately does not do.
SERVED_PAIR_TYPES = ("similar", "complement")


def servable_neighbours(tenant_id: str, product_key: str, limit: int = 3) -> list:
    """Whole product rows a shopper may be shown for this anchor.

    Stock and missing_fields are filtered HERE rather than trusted from pairing
    time, because a product can sell out an hour after the graph was built.
    """
    conn = get_db_connection()
    try:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            cur.execute(sql.SQL("""
                SELECT p.*, n.pair_type, n.score AS pair_score,
                       n.confidence AS pair_confidence, n.source AS pair_source,
                       n.neighbor_key, n.anchor_key
                FROM {}.strategist_product_neighbors n
                JOIN {}.strategist_products p ON p.product_key = n.neighbor_key
                WHERE n.anchor_key = %s
                  AND n.pair_type = ANY(%s)
                  AND p.in_stock
                  -- Only a blocking gap hides a product from shoppers. A
                  -- missing reference price is not one: it affects
                  -- cross-currency comparison, not whether the card can be
                  -- shown or clicked.
                  AND NOT (p.missing_fields && %s)
                ORDER BY n.score DESC, n.neighbor_key
            """).format(sql.Identifier(tenant_id), sql.Identifier(tenant_id)),
                (product_key, list(SERVED_PAIR_TYPES), list(BLOCKING_FIELDS)))
            rows = [dict(r) for r in cur.fetchall()]
            # Same connection, and scoped to this anchor: opening a connection
            # costs ~200ms and this runs on every page view of every visitor.
            # is_servable stays the authority on what may be shown -- only the
            # fetch is narrowed.
            decisions = load_decisions_for_anchor(tenant_id, product_key, conn)
    finally:
        conn.close()

    servable = []
    for row in rows:
        pair = {"confidence": row["pair_confidence"], "source": row["pair_source"]}
        decision = decisions.get(
            (row["anchor_key"], row["neighbor_key"], row["pair_type"]))
        if not is_servable(pair, decision):
            continue
        servable.append(row)
        if len(servable) >= limit:
            break
    return servable

# The bands the review screen offers: "80% and above", "60-79%", "under 60%".
# On the live catalogue they split 1,117 pairs into 356 / 403 / 358, so each
# one is a pile a merchant can actually work through.
SCORE_BANDS = {
    "high": (0.80, 1.01),
    "medium": (0.60, 0.80),
    "low": (0.0, 0.60),
}

DECISION_STATES = ("pending", "approved", "rejected", "auto", "all")



def _state_of(row: dict, decision) -> str:
    """What a merchant would call this pair's state."""
    if decision == "approved":
        return "approved"
    if decision == "rejected":
        return "rejected"
    # No human has ruled on it. Either it cleared the confidence threshold and
    # is already being served, or it is waiting for someone.
    return "auto" if is_servable(row, None) else "pending"


def _matches(row: dict, state: str, status: str, category, score) -> bool:
    if status != "all" and state != status:
        return False
    if category:
        # _card() pops the flat columns into a nested card, so this has to read
        # whichever shape the row is in: list_pairings filters before carding,
        # list_anchors re-filters rows that have already been carded.
        anchor_category = (row.get("anchor_category")
                           or (row.get("anchor") or {}).get("category"))
        if anchor_category != category:
            return False
    if score:
        low, high = SCORE_BANDS[score]
        if not low <= row["score"] < high:
            return False
    return True


def list_pairings(tenant_id: str, limit: int = 100, anchor_key: str = None,
                  status: str = "pending", category: str = None,
                  score: str = None) -> dict:
    """Pairs, filtered the way a review screen wants them.

    One endpoint rather than one per view: a merchant working through matches
    wants to see what is waiting, then what they approved, then only the strong
    ones in a category, and swapping endpoints for each of those would push the
    same filtering into every client.
    """
    if status not in DECISION_STATES:
        raise ValueError(f"unknown status: {status}; expected one of {DECISION_STATES}")
    if score and score not in SCORE_BANDS:
        raise ValueError(f"unknown score band: {score}; expected one of {tuple(SCORE_BANDS)}")

    conn = get_db_connection()
    try:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            cur.execute(sql.SQL(
                "SELECT COUNT(*) AS total FROM {}.strategist_product_neighbors"
            ).format(sql.Identifier(tenant_id)))
            total = cur.fetchone()["total"]

            cur.execute(sql.SQL("""
                SELECT n.anchor_key, n.neighbor_key, n.pair_type, n.score,
                       n.confidence, n.source, n.reasons, n.computed_at,
                       a.name AS anchor_name, a.category AS anchor_category,
                       a.price_cents AS anchor_price_cents,
                       a.currency AS anchor_currency,
                       a.image_url AS anchor_image_url,
                       a.in_stock AS anchor_in_stock,
                       p.name AS neighbor_name, p.category AS neighbor_category,
                       p.price_cents AS neighbor_price_cents,
                       p.currency AS neighbor_currency,
                       p.image_url AS neighbor_image_url,
                       p.in_stock AS neighbor_in_stock
                FROM {}.strategist_product_neighbors n
                JOIN {}.strategist_products a ON a.product_key = n.anchor_key
                JOIN {}.strategist_products p ON p.product_key = n.neighbor_key
                ORDER BY n.score DESC
            """).format(sql.Identifier(tenant_id), sql.Identifier(tenant_id),
                        sql.Identifier(tenant_id)))
            rows = [dict(r) for r in cur.fetchall()]
    finally:
        conn.close()

    # Counted before filtering, so it always reports every pair whose product
    # has vanished since the last run, not just the ones this filter would show.
    dropped = total - len(rows)

    decisions = load_decisions(tenant_id)
    matched = []
    for row in rows:
        decision = decisions.get(
            (row["anchor_key"], row["neighbor_key"], row["pair_type"]))
        state = _state_of(row, decision)
        if anchor_key and row["anchor_key"] != anchor_key:
            continue
        if not _matches(row, state, status, category, score):
            continue
        row["state"] = state
        row["decision"] = decision
        row["anchor"] = _card(row, "anchor")
        row["neighbor"] = _card(row, "neighbor")
        matched.append(row)
        if len(matched) >= limit:
            break

    return {"pairings": matched, "dropped": dropped}


def list_anchors(tenant_id: str, limit: int = 50, status: str = "pending",
                 category: str = None, score: str = None) -> dict:
    """Products with matches, and how many of each state they hold.

    A merchant thinks in products, not in pairs: "this dress has 4 to review"
    is reviewable, a flat list of 229 pairs from 80 different products is not.

    Every anchor carries the full breakdown, not just the filtered count, so a
    screen can show "4 pending, 2 approved" without asking again per product.
    """
    # Validated here too: this reads everything and filters in Python, so an
    # unknown status would otherwise match nothing and return an empty list --
    # a typo would look exactly like "you have no work to do".
    if status not in DECISION_STATES:
        raise ValueError(f"unknown status: {status}; expected one of {DECISION_STATES}")
    if score and score not in SCORE_BANDS:
        raise ValueError(f"unknown score band: {score}; expected one of {tuple(SCORE_BANDS)}")

    everything = list_pairings(tenant_id, limit=10_000, status="all")["pairings"]

    grouped = {}
    for row in everything:
        entry = grouped.setdefault(row["anchor_key"], {
            "anchor_key": row["anchor_key"],
            "anchor": row["anchor"],
            "matching": 0,
            "pending": 0, "approved": 0, "rejected": 0, "auto": 0,
            "top_score": 0.0,
            "avg_score": 0.0,
            "_scores": [],
        })
        entry[row["state"]] += 1
        if _matches(row, row["state"], status, category, score):
            entry["matching"] += 1
            entry["top_score"] = max(entry["top_score"], row["score"])
            entry["_scores"].append(row["score"])

    anchors = []
    for entry in grouped.values():
        if not entry["matching"]:
            continue
        scores = entry.pop("_scores")
        # The screen shows an average match alongside the count.
        entry["avg_score"] = round(sum(scores) / len(scores), 4)
        anchors.append(entry)
    # Best candidate first: a merchant working down the list should meet the
    # most promising decisions while they still have patience for them.
    anchors.sort(key=lambda a: (-a["top_score"], a["anchor_key"]))
    shown = anchors[:limit]
    all_scores = [s for a in shown for s in [a["avg_score"]]]
    return {
        "anchors": shown,
        "total": len(anchors),
        "avg_score": round(sum(all_scores) / len(all_scores), 4) if all_scores else 0.0,
    }
