"""Embeds products once and caches the vectors on content_hash.

Phase 1 proved this mechanism: its second run over 218 products made zero model
calls. Re-pairing after a small catalog change must re-embed only what changed.
"""
import logging

from psycopg2 import sql

from app.services.infra.database import get_db_connection
from app.services.enrichment.job import NO_CONTENT_HASH
from app.services.catalog.products import _cosine, _embed_batch

logger = logging.getLogger(__name__)


def cosine(a: list, b: list) -> float:
    return _cosine(a, b)


def embedding_text(product: dict) -> str:
    """Sorted, so attribute order -- which is not stable across runs -- cannot
    change the text and force a needless re-embed."""
    attributes = sorted(
        f"{a.get('key')} {a.get('value')}"
        for a in (product.get("attributes") or [])
        if a.get("key") and a.get("value") is not None)
    parts = [product.get("name") or "", product.get("description") or "",
             product.get("category") or "", " ".join(attributes)]
    return ". ".join(p.strip() for p in parts if p and p.strip())


def load_cached(tenant_id: str) -> dict:
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(sql.SQL(
                "SELECT product_key, content_hash, vector "
                "FROM {}.strategist_product_embeddings"
            ).format(sql.Identifier(tenant_id)))
            return {key: (hash_, list(vector))
                    for key, hash_, vector in cur.fetchall()}
    finally:
        conn.close()


def store(tenant_id: str, product_key: str, content_hash, vector: list) -> None:
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(sql.SQL(
                "INSERT INTO {}.strategist_product_embeddings "
                "(product_key, content_hash, vector) VALUES (%s, %s, %s) "
                "ON CONFLICT (product_key) DO UPDATE SET "
                "content_hash = EXCLUDED.content_hash, "
                "vector = EXCLUDED.vector, computed_at = CURRENT_TIMESTAMP"
            ).format(sql.Identifier(tenant_id)),
                (product_key, content_hash or NO_CONTENT_HASH, vector))
        conn.commit()
    except Exception:
        conn.rollback()
        logger.error("Could not cache an embedding for %s", product_key,
                     exc_info=True)
    finally:
        conn.close()


def vectors_for(tenant_id: str, products: list) -> dict:
    cached = load_cached(tenant_id)
    vectors, stale = {}, []

    for product in products:
        key = product["product_key"]
        current = product.get("content_hash") or NO_CONTENT_HASH
        entry = cached.get(key)
        # COALESCE-style comparison on both sides: crawled rows carry no
        # content_hash, and a raw NULL comparison would re-embed them every run.
        if entry and entry[0] == current and entry[1]:
            vectors[key] = entry[1]
        else:
            stale.append(product)

    if stale:
        computed = _embed_batch([embedding_text(p) for p in stale],
                                tenant_id=tenant_id)
        if len(computed) != len(stale):
            logger.error("Embedding failed for %s; %d products will be scored "
                         "on attributes alone", tenant_id, len(stale))
            return vectors
        for product, vector in zip(stale, computed):
            vectors[product["product_key"]] = vector
            store(tenant_id, product["product_key"],
                  product.get("content_hash"), vector)

    logger.info("%s: %d embeddings reused, %d computed",
                tenant_id, len(products) - len(stale), len(stale))
    return vectors
