"""Merchant approve/reject decisions on candidate product pairs.

The only writer of strategist_pairing_decisions. The pairing job rewrites its
own table (strategist_product_neighbors) on every run; if decisions lived
there, a merchant who rejected a pair would see it come back after the next
catalogue sync. Keeping decisions in a separate table the job never touches
is what makes a re-run safe.
"""
import logging

from psycopg2 import sql
from psycopg2.extras import execute_values

from app.services.infra.database import get_db_connection

logger = logging.getLogger(__name__)

# Measured against the live 218-product catalog: confidence is not a continuous
# distribution but flat spikes per type -- complement 0.75, upsell 0.95 -- plus
# similar spread over 0.55..0.90. The threshold is therefore a switch over which
# types queue, rather than a cut through a distribution.
#
# At 0.55 (SIMILAR_MIN_COSINE) everything the rules emit is servable and the
# queue is empty, which makes the approval feature inert. 0.70 queues the weaker
# similar pairs -- the least certain thing being served -- and leaves complements
# and upsells to go straight out. On the live catalogue that is a few pairs each
# across a few dozen products, which is a queue a merchant can actually finish.
#
# Revisit once several merchants' catalogues give confidence a real spread.
APPROVAL_THRESHOLD = 0.70

_VALID_DECISIONS = {"approved", "rejected"}


def record(tenant_id: str, decisions: list) -> int:
    """Upserts a batch of merchant decisions. Returns the number recorded."""
    for entry in decisions:
        if entry.get("decision") not in _VALID_DECISIONS:
            raise ValueError(
                f"Invalid decision {entry.get('decision')!r}; must be one of "
                f"{sorted(_VALID_DECISIONS)}")

    rows = [
        (entry["anchor_key"], entry["neighbor_key"], entry["pair_type"],
         entry["decision"], entry.get("decided_by"))
        for entry in decisions
    ]

    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            execute_values(cur, sql.SQL("""
                INSERT INTO {}.strategist_pairing_decisions
                    (anchor_key, neighbor_key, pair_type, decision, decided_by)
                VALUES %s
                ON CONFLICT (anchor_key, neighbor_key, pair_type) DO UPDATE
                SET decision = EXCLUDED.decision,
                    decided_by = EXCLUDED.decided_by,
                    decided_at = CURRENT_TIMESTAMP
            """).format(sql.Identifier(tenant_id)), rows)
            conn.commit()
            return len(rows)
    except Exception:
        conn.rollback()
        logger.error("Failed to record pairing decisions for %s", tenant_id, exc_info=True)
        raise
    finally:
        conn.close()


def load_for_anchor(tenant_id: str, anchor_key: str, conn) -> dict:
    """One anchor's decisions, on a caller-supplied connection.

    The serving path runs on every page view and opening a connection costs
    ~200ms here, so it cannot afford its own; nor can it afford to read every
    decision the tenant has ever made to answer about one product.
    """
    with conn.cursor() as cur:
        cur.execute(sql.SQL(
            "SELECT anchor_key, neighbor_key, pair_type, decision "
            "FROM {}.strategist_pairing_decisions WHERE anchor_key = %s"
        ).format(sql.Identifier(tenant_id)), (anchor_key,))
        return {(row[0], row[1], row[2]): row[3] for row in cur.fetchall()}


def load(tenant_id: str) -> dict:
    """All of a tenant's decisions, keyed (anchor_key, neighbor_key, pair_type).

    A pairing run needs the full set to decide servability for every
    candidate it produces, so one query beats one lookup per pair.
    """
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(sql.SQL(
                "SELECT anchor_key, neighbor_key, pair_type, decision "
                "FROM {}.strategist_pairing_decisions"
            ).format(sql.Identifier(tenant_id)))
            return {(row[0], row[1], row[2]): row[3] for row in cur.fetchall()}
    finally:
        conn.close()


def is_servable(pair: dict, decision) -> bool:
    """The §5 servability rule: a merchant's decision always wins; absent one,
    a merchant-declared pair is trusted outright and everything else needs
    enough confidence to skip the approval queue.
    """
    if decision == "rejected":
        return False
    if decision == "approved":
        return True
    # A merchant's own declaration outranks any inference, so it never queues.
    if pair.get("source") == "merchant":
        return True
    return (pair.get("confidence") or 0.0) >= APPROVAL_THRESHOLD
