"""Builds a tenant's pairing graph.

The job owns strategist_product_neighbors and rewrites it freely. It never writes
to strategist_pairing_decisions -- a merchant who rejects a pair and then
re-syncs must not see that pair return.
"""
import logging

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

from app.services.infra.database import get_db_connection, migrate_pairing_tables
from app.services.pairing.candidates import candidate_pairs, load_eligible
from app.services.pairing.decisions import is_servable, load as load_decisions
from app.services.pairing.embeddings import cosine, vectors_for
from app.services.pairing.rules import (
    PAIR_TYPES, score_complement, score_similar, score_upsell,
)

logger = logging.getLogger(__name__)

MAX_NEIGHBORS_PER_TYPE = 8


def score_pair(a: dict, b: dict, vectors: dict) -> list:
    left, right = vectors.get(a["product_key"]), vectors.get(b["product_key"])
    # An embedding failure degrades a pair to attribute-only scoring rather than
    # aborting the run: a partial graph beats no graph.
    similarity = cosine(left, right) if left and right else 0.0

    results = []
    for pair_type, scored in (
        ("similar", score_similar(a, b, similarity)),
        ("complement", score_complement(a, b, similarity)),
        ("upsell", score_upsell(a, b)),
    ):
        if scored:
            results.append({"anchor_key": a["product_key"],
                            "neighbor_key": b["product_key"],
                            "pair_type": pair_type, **scored})
    return results


def merchant_pairs(products: list) -> list:
    """A merchant's own declarations outrank every inference in this module."""
    known = {p["product_key"] for p in products}
    pairs = []
    for product in products:
        declared = (product.get("tenant_relations") or {}).get("complementary") or []
        for neighbor in declared:
            if neighbor not in known:
                continue
            pairs.append({"anchor_key": product["product_key"],
                          "neighbor_key": neighbor, "pair_type": "complement",
                          "score": 1.0, "confidence": 1.0, "source": "merchant",
                          "reasons": ["declared by the merchant"]})
    return pairs


def _top_per_anchor_and_type(pairs: list) -> list:
    grouped = {}
    for pair in pairs:
        grouped.setdefault((pair["anchor_key"], pair["pair_type"]), []).append(pair)

    kept = []
    for members in grouped.values():
        members.sort(key=lambda p: p["score"], reverse=True)
        kept.extend(members[:MAX_NEIGHBORS_PER_TYPE])
    return kept


def replace_neighbors(tenant_id: str, pairs: list) -> None:
    """Replaces this tenant's whole graph in one transaction.

    Delete-then-insert rather than upsert: a pair that no longer scores must
    disappear, and an upsert would leave last run's edges behind forever.
    """
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(sql.SQL(
                "DELETE FROM {}.strategist_product_neighbors"
            ).format(sql.Identifier(tenant_id)))
            if pairs:
                # execute_values already handles a Composable query (it calls
                # .as_string(cur) on it internally), so passing the composed
                # sql.SQL object straight through is both correct and simpler
                # than pre-stringifying it ourselves.
                execute_values(cur, sql.SQL(
                    "INSERT INTO {}.strategist_product_neighbors "
                    "(anchor_key, neighbor_key, pair_type, score, confidence, "
                    "source, reasons) VALUES %s"
                ).format(sql.Identifier(tenant_id)),
                    [(p["anchor_key"], p["neighbor_key"], p["pair_type"],
                      p["score"], p["confidence"], p["source"],
                      Json(p["reasons"])) for p in pairs])
        conn.commit()
    except Exception:
        conn.rollback()
        logger.error("Could not write the pairing graph for %s", tenant_id,
                     exc_info=True)
        raise
    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 pair_tenant(tenant_id: str, progress=None) -> dict:
    migrate_pairing_tables(tenant_id)
    products = load_eligible(tenant_id)
    _report(progress, 5, "loading eligible products")
    if len(products) < 2:
        _report(progress, 100, "no eligible pairs to score")
        return {"products": len(products), "pairs": 0, "servable": 0,
                "queued": 0, "by_type": {}}

    vectors = vectors_for(tenant_id, products)
    _report(progress, 45, "embedding products")

    candidates = candidate_pairs(products)
    total_candidates = len(candidates)
    scored = []
    # Report only when the percentage actually changes. The live catalogue
    # produces 9,187 candidate pairs across 45 distinct percentages, and a tick
    # per pair would be 9,187 Redis round trips to say the same 45 things.
    last_percent = 45
    for i, (a, b) in enumerate(candidates, start=1):
        scored.extend(score_pair(a, b, vectors))
        percent = 45 + (45 * i // total_candidates)
        if percent != last_percent:
            last_percent = percent
            _report(progress, percent,
                    f"scoring candidate pairs — {i} of {total_candidates}")
    if not total_candidates:
        _report(progress, 90, "scoring candidate pairs")

    # Merchant declarations are added last and deduplicated in favour of
    # themselves: their confidence of 1.0 must never be lowered by an inferred
    # duplicate of the same edge.
    declared = merchant_pairs(products)
    declared_keys = {(p["anchor_key"], p["neighbor_key"], p["pair_type"])
                     for p in declared}
    scored = [p for p in scored
              if (p["anchor_key"], p["neighbor_key"], p["pair_type"])
              not in declared_keys] + declared

    kept = _top_per_anchor_and_type(scored)
    replace_neighbors(tenant_id, kept)
    _report(progress, 100, "writing pairing graph")

    decisions = load_decisions(tenant_id)
    servable = sum(1 for p in kept if is_servable(
        p, decisions.get((p["anchor_key"], p["neighbor_key"], p["pair_type"]))))

    by_type = {t: sum(1 for p in kept if p["pair_type"] == t) for t in PAIR_TYPES}
    report = {"products": len(products), "pairs": len(kept),
              "servable": servable, "queued": len(kept) - servable,
              "by_type": by_type}
    logger.info("Paired %s: %s", tenant_id, report)
    return report
