"""Storing and querying the tenant's products.

No parsing here -- that is product_extraction. This module owns the table, the
URL matching rule, and the shape products take when they leave the backend.
"""
import json
import logging
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse

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

from app.core.config import settings
from app.core.llm_client import client
from app.services.infra.database import get_db_connection, insert_llm_usage, search_vector_data
from app.services.infra.embedding import get_embeddings
from app.services.infra.quotas import check_quota, estimate_text_tokens

logger = logging.getLogger(__name__)

# Parameters that identify a campaign, not a product.
_TRACKING_PARAMS = {"utm_source", "utm_medium", "utm_campaign", "utm_term",
                    "utm_content", "utm_id", "gclid", "fbclid", "msclkid",
                    "mc_cid", "mc_eid", "ref", "referrer"}

# How many texts go into a single embeddings API call.
_EMBED_BATCH_SIZE = 64

# Relatedness is computed for at most this many products (the newest ones), so
# ingesting a large catalogue cannot turn into thousands of sequential
# embedding calls inside a single request.
MAX_RELATED_PRODUCTS = 500


def normalise_url(url: str) -> str:
    """Make two spellings of the same page match.

    Strips the fragment, a trailing slash, and known tracking parameters. Keeps
    every other query parameter: plenty of sites identify the product there
    (?product_id=123), and dropping the query wholesale would collapse a whole
    catalogue onto one URL.

    Tolerant of non-string input (e.g. None): callers that pass a bad value get
    "" back rather than an AttributeError, so lookups return empty rather than
    raise.
    """
    url = str(url or "").strip()
    if not url:
        return ""
    parts = urlparse(url)
    query = [(k, v) for k, v in parse_qsl(parts.query, keep_blank_values=True)
             if k.lower() not in _TRACKING_PARAMS]
    path = parts.path.rstrip("/") or "/"
    return urlunparse((parts.scheme, parts.netloc, path, "", urlencode(query), ""))


def _table_exists(conn, tenant_id: str) -> bool:
    """Whether the tenant has a products table.

    Opens its own plain cursor rather than taking the caller's, because
    callers use different cursor factories (`save_products` and
    `save_related` use a plain cursor, `_query` uses a `RealDictCursor`) and a
    row read as `row[0]` breaks under a dict cursor while `row["exists"]`
    breaks under a plain one. A private cursor keeps this function correct
    under either.
    """
    with conn.cursor() as cur:
        cur.execute(
            "SELECT EXISTS (SELECT FROM information_schema.tables "
            "WHERE table_schema = %s AND table_name = 'strategist_products')",
            (tenant_id,))
        return cur.fetchone()[0]


def save_products(tenant_id: str, products: list, crawled_urls: list,
                  crawl_complete: bool) -> dict:
    """Upsert products, and delete stale ones only when it is safe to.

    A crawl that failed halfway, or that was truncated by max_pages, would
    otherwise delete products that are still on sale. When `crawl_complete` is
    false, this writes and updates but deletes nothing.
    """
    saved = deleted = 0
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            if not _table_exists(conn, tenant_id):
                logger.warning(f"No products table for {tenant_id}; skipping save.")
                return {"saved": 0, "deleted": 0}

            for product in products:
                cur.execute(sql.SQL("""
                    INSERT INTO {}.strategist_products
                        (product_key, name, description, image_url, product_url,
                         category, ctas, options, raw, extracted_at)
                    VALUES (%s,%s,%s,%s,%s,%s,%s,%s,%s, CURRENT_TIMESTAMP)
                    ON CONFLICT (product_key) DO UPDATE SET
                        name = EXCLUDED.name,
                        description = EXCLUDED.description,
                        image_url = EXCLUDED.image_url,
                        product_url = EXCLUDED.product_url,
                        category = EXCLUDED.category,
                        ctas = EXCLUDED.ctas,
                        options = EXCLUDED.options,
                        raw = EXCLUDED.raw,
                        extracted_at = CURRENT_TIMESTAMP
                """).format(sql.Identifier(tenant_id)), (
                    product["product_key"], product["name"], product.get("description"),
                    product.get("image_url"), normalise_url(product["product_url"]),
                    product.get("category"), Json(product.get("ctas") or []),
                    Json(product.get("options") or []), Json(product.get("raw") or {}),
                ))
                saved += 1

            if crawl_complete and crawled_urls:
                # The crawler, Shopify and the HTTP API all share this table, so
                # this delete must only ever reap rows the crawler itself owns --
                # otherwise a website re-crawl wipes out the whole synced catalog.
                cur.execute(sql.SQL(
                    "DELETE FROM {}.strategist_products "
                    "WHERE source_kind = 'crawl' AND product_url <> ALL(%s)"
                ).format(sql.Identifier(tenant_id)),
                    ([normalise_url(u) for u in crawled_urls],))
                deleted = cur.rowcount
        conn.commit()
        logger.info(f"Products for {tenant_id}: {saved} saved, {deleted} deleted "
                    f"(crawl_complete={crawl_complete})")
        return {"saved": saved, "deleted": deleted}
    except Exception as ex:
        conn.rollback()
        logger.error(f"Could not save products for {tenant_id}: {ex}", exc_info=True)
        return {"saved": 0, "deleted": 0}
    finally:
        conn.close()


def delete_product(tenant_id: str, product_key: str) -> bool:
    """Remove one product and everything that references it -- its pairings
    (as either anchor or neighbor) and its embedding. Returns False if the
    key didn't exist, so the caller can 404 rather than report a silent
    no-op as success.
    """
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(sql.SQL(
                "DELETE FROM {}.strategist_product_neighbors "
                "WHERE anchor_key = %s OR neighbor_key = %s"
            ).format(sql.Identifier(tenant_id)), (product_key, product_key))
            cur.execute(sql.SQL(
                "DELETE FROM {}.strategist_product_embeddings WHERE product_key = %s"
            ).format(sql.Identifier(tenant_id)), (product_key,))
            cur.execute(sql.SQL(
                "DELETE FROM {}.strategist_products WHERE product_key = %s"
            ).format(sql.Identifier(tenant_id)), (product_key,))
            removed = cur.rowcount > 0
        conn.commit()
        return removed
    except Exception:
        conn.rollback()
        logger.error(f"Could not delete product {product_key} for {tenant_id}",
                    exc_info=True)
        raise
    finally:
        conn.close()


def delete_all_products(tenant_id: str) -> dict:
    """Wipe every product, pairing, and embedding for a tenant.

    Connected sources are the caller's concern, not this function's -- see
    app.api.catalog's delete_all_products_endpoint, which also disconnects
    them. A source left connected after this would just get silently
    re-crawled on the next build, repopulating exactly what was deleted."""
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            cur.execute(sql.SQL(
                "DELETE FROM {}.strategist_product_neighbors"
            ).format(sql.Identifier(tenant_id)))
            pairings = cur.rowcount
            cur.execute(sql.SQL(
                "DELETE FROM {}.strategist_product_embeddings"
            ).format(sql.Identifier(tenant_id)))
            embeddings = cur.rowcount
            cur.execute(sql.SQL(
                "DELETE FROM {}.strategist_products"
            ).format(sql.Identifier(tenant_id)))
            products = cur.rowcount
        conn.commit()
        logger.info(f"Wiped catalogue for {tenant_id}: {products} products, "
                    f"{pairings} pairings, {embeddings} embeddings")
        return {"products": products, "pairings": pairings, "embeddings": embeddings}
    except Exception:
        conn.rollback()
        logger.error(f"Could not wipe catalogue for {tenant_id}", exc_info=True)
        raise
    finally:
        conn.close()


def _query(tenant_id: str, where: str, params: tuple, limit: int) -> list:
    conn = get_db_connection()
    try:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            if not _table_exists(conn, tenant_id):
                return []
            cur.execute(sql.SQL(
                "SELECT * FROM {}.strategist_products WHERE " + where +
                " ORDER BY extracted_at DESC LIMIT %s"
            ).format(sql.Identifier(tenant_id)), params + (limit,))
            return [dict(r) for r in cur.fetchall()]
    except Exception as ex:
        logger.error(f"Product query failed for {tenant_id}: {ex}", exc_info=True)
        return []
    finally:
        conn.close()


def find_by_url(tenant_id: str, url: str) -> dict:
    rows = _query(tenant_id, "product_url = %s", (normalise_url(url),), 1)
    return rows[0] if rows else {}


def find_by_urls(tenant_id: str, urls: list) -> list:
    if not urls:
        return []
    return _query(tenant_id, "product_url = ANY(%s)",
                  ([normalise_url(u) for u in urls],), len(urls))


def list_products(tenant_id: str, category: str = None, limit: int = 3) -> list:
    where, params = "TRUE", ()
    if category:
        where += " AND lower(category) = lower(%s)"
        params += (category,)
    return _query(tenant_id, where, params, limit)


def find_by_keys(tenant_id: str, keys: list) -> list:
    if not keys:
        return []
    return _query(tenant_id, "product_key = ANY(%s)", (list(keys),), len(keys))


def search_products(tenant_id: str, query: str, limit: int = 3) -> list:
    """Products matching what the visitor actually meant.

    The product pages are already chunked and embedded in the knowledge base, so
    the search happens there and the winning chunks' source_url is mapped back
    onto products. This is a semantic match: "something light for jogging" finds
    a running shoe whose description uses neither word, where a category or
    keyword filter would return nothing.
    """
    query = (query or "").strip()
    if not query:
        return []
    try:
        embedding = get_embeddings(query, tenant_id=tenant_id)
        hits = search_vector_data(tenant_id, embedding, limit=12)
    except Exception as ex:
        logger.error(f"Product search failed for {tenant_id}: {ex}", exc_info=True)
        return []

    # Preserve retrieval order: the first URL matched is the best match.
    urls, seen = [], set()
    for hit in hits or []:
        url = normalise_url((hit.get("metadata") or {}).get("source_url") or "")
        if url and url not in seen:
            seen.add(url)
            urls.append(url)

    rows = {r["product_url"]: r for r in find_by_urls(tenant_id, urls)}
    return [rows[u] for u in urls if u in rows][:limit]


# Below this, a product is not really an answer to the question. Cosine over
# this embedding model puts unrelated products around 0.1-0.2, so the floor is
# set to exclude those while keeping a loose but genuine match: "something for
# jogging" should still reach a running shoe.
SEMANTIC_FLOOR = 0.25


def semantic_products(tenant_id: str, query: str, limit: int = 3) -> list:
    """Products closest in meaning to what the visitor asked for.

    Searches the products' own embeddings, which the pairing build writes for
    every product whatever its source. That is the difference from
    search_products, which searches the knowledge base and can therefore only
    ever reach products whose page was crawled -- on a live catalogue of 218
    products from three sources, it returned a hoodie for "snowboard" because
    the scraped pages were the only ones it could see.

    Understands the question rather than matching its words: "something light
    for jogging" reaches a running shoe whose description contains neither.

    Costs one embedding call for the query. Fine in the chat, where a model
    call is already in the loop; not fine on a path that runs per page view.
    """
    query = (query or "").strip()
    if not query:
        return []

    # Imported here: pairing.embeddings already imports _cosine from this
    # module, so taking its loader at module scope would be a cycle.
    from app.services.pairing.embeddings import load_cached

    try:
        cached = load_cached(tenant_id)
        if not cached:
            return []
        query_vector = get_embeddings(query, tenant_id=tenant_id)
    except Exception as ex:
        logger.error(f"Semantic product search failed for {tenant_id}: {ex}",
                     exc_info=True)
        return []

    scored = []
    for product_key, (_, vector) in cached.items():
        if not vector:
            continue
        score = _cosine(query_vector, vector)
        if score >= SEMANTIC_FLOOR:
            scored.append((score, product_key))

    if not scored:
        return []

    scored.sort(reverse=True)
    keys = [key for _, key in scored[:limit]]
    # find_by_keys returns rows in the table's order, not the ranking's.
    rows = {r["product_key"]: r for r in find_by_keys(tenant_id, keys)}
    return [rows[k] for k in keys if k in rows]


def match_products(tenant_id: str, query: str, limit: int = 3) -> list:
    """Products whose text plainly contains the visitor's search term.

    A single ILIKE query against name, description and category, ordered so a
    name match outranks a description match and a description match outranks a
    category match. No embedding call and no quota check: unlike
    `search_products`, this is meant to run on every page view of every
    visitor, so it has to be a deterministic database lookup rather than a
    model round-trip. Never raises -- `_query` already returns [] on failure.
    """
    query = (query or "").strip()
    if not query:
        return []
    conn = get_db_connection()
    try:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            if not _table_exists(conn, tenant_id):
                return []
            like = f"%{query}%"
            cur.execute(sql.SQL("""
                SELECT *, CASE
                    WHEN name ILIKE %s THEN 0
                    WHEN description ILIKE %s THEN 1
                    ELSE 2
                END AS _rank
                FROM {}.strategist_products
                WHERE name ILIKE %s OR description ILIKE %s OR category ILIKE %s
                ORDER BY _rank, extracted_at DESC
                LIMIT %s
            """).format(sql.Identifier(tenant_id)),
                (like, like, like, like, like, limit))
            return [dict(r) for r in cur.fetchall()]
    except Exception as ex:
        logger.error(f"Product text match failed for {tenant_id}: {ex}", exc_info=True)
        return []
    finally:
        conn.close()


def related_products(tenant_id: str, row: dict, limit: int = 3) -> list:
    """A product's nearest neighbours, computed at ingestion.

    A lookup rather than a similarity search, because this runs on every page
    view and the thinking has already been done.
    """
    keys = list(row.get("related_keys") or [])[:limit]
    if not keys:
        return []
    found = {r["product_key"]: r for r in find_by_keys(tenant_id, keys)}
    return [found[k] for k in keys if k in found]


def _cosine(a: list, b: list) -> float:
    """Similarity between two embeddings, 0.0 when either is degenerate."""
    dot = sum(x * y for x, y in zip(a, b))
    na = sum(x * x for x in a) ** 0.5
    nb = sum(y * y for y in b) ** 0.5
    return dot / (na * nb) if na and nb else 0.0


def _embed_batch(texts: list, tenant_id: str = None, user_id: str = None) -> list:
    """Embed many texts in as few API calls as possible, in input order.

    One call per _EMBED_BATCH_SIZE texts instead of one call per text, so
    computing relatedness for a real catalogue is dozens of round-trips rather
    than thousands. Usage is logged once per batch, the same way
    `get_embeddings` logs once per call. Never raises: any failure returns []
    and the caller treats that as "skip relatedness this run".
    """
    vectors = []
    try:
        for start in range(0, len(texts), _EMBED_BATCH_SIZE):
            batch = texts[start:start + _EMBED_BATCH_SIZE]
            if tenant_id:
                requested = sum(estimate_text_tokens(t) for t in batch)
                check_quota(tenant_id, "ai_tokens", user_id=user_id, requested_amount=requested)

            response = client.embeddings.create(
                input=batch, model=settings.EMBED_MODEL, timeout=30.0)
            usage = response.usage
            logger.info(f"LLM Usage (Embedding batch of {len(batch)}): "
                        f"Prompt: {usage.prompt_tokens}, Total: {usage.total_tokens}")
            if tenant_id:
                insert_llm_usage(tenant_id, "Embedding", settings.EMBED_MODEL,
                                 usage.prompt_tokens, 0, usage.total_tokens)
            vectors.extend(item.embedding for item in response.data)
        return vectors
    except Exception as ex:
        logger.error(f"Batch embedding failed for {tenant_id}: {ex}", exc_info=True)
        return []


def compute_relatedness(tenant_id: str, limit: int = 5, user_id: str = None) -> int:
    """Work out each product's nearest neighbours, once, at ingestion.

    Category matching is a reflex: it cannot tell that a running shoe goes with
    running socks but not a running-themed poster, and it fails outright on the
    many products that carry no category at all. Comparing what the products
    actually say does better.

    Doing it here rather than per request is what lets the proactive endpoint stay
    a fast lookup while still being a real recommendation. Never raises.

    Capped at MAX_RELATED_PRODUCTS (the newest rows -- `_query` already orders
    by extracted_at DESC): a catalogue larger than that would otherwise turn a
    single ingestion request into thousands of sequential embedding calls and
    an O(n^2) comparison. Skipping the rest is logged, not silent.
    """
    rows = _query(tenant_id, "TRUE", (), 10000)
    if len(rows) < 2:
        return 0

    if len(rows) > MAX_RELATED_PRODUCTS:
        skipped = len(rows) - MAX_RELATED_PRODUCTS
        rows = rows[:MAX_RELATED_PRODUCTS]
        logger.warning(f"{tenant_id} has {skipped} products beyond the "
                       f"{MAX_RELATED_PRODUCTS}-product relatedness cap; skipping them.")

    texts = [f"{row['name']}. {row.get('description') or ''}".strip() for row in rows]
    vectors = _embed_batch(texts, tenant_id=tenant_id, user_id=user_id)
    if len(vectors) != len(rows):
        logger.error(f"Could not embed products for {tenant_id}")
        return 0

    updated = 0
    for i, row in enumerate(rows):
        scored = sorted(
            ((_cosine(vectors[i], vectors[j]), rows[j]["product_key"])
             for j in range(len(rows)) if j != i),
            reverse=True)
        save_related(tenant_id, row["product_key"], [k for _, k in scored[:limit]])
        updated += 1

    logger.info(f"Computed relatedness for {updated} products in {tenant_id}")
    return updated


def save_related(tenant_id: str, product_key: str, related: list) -> None:
    """Store a product's nearest neighbours. Never raises."""
    conn = get_db_connection()
    try:
        with conn.cursor() as cur:
            if not _table_exists(conn, tenant_id):
                return
            cur.execute(sql.SQL(
                "UPDATE {}.strategist_products SET related_keys = %s WHERE product_key = %s"
            ).format(sql.Identifier(tenant_id)), (list(related), product_key))
        conn.commit()
    except Exception as ex:
        conn.rollback()
        logger.error(f"Could not save related products for {product_key}: {ex}")
    finally:
        conn.close()


def to_card(row: dict) -> dict:
    """The shape a product takes when it leaves the backend.

    Carries no price, currency or stock level: those go stale fastest between
    crawls, and a wrong price on a card shown to a customer is the worst thing
    this feature could do. The CTA takes the visitor to the page, where those
    numbers are correct by definition.
    """
    ctas = row.get("ctas") or []
    if isinstance(ctas, str):
        ctas = json.loads(ctas)
    options = row.get("options") or []
    if isinstance(options, str):
        options = json.loads(options)
    return {
        "product_id": row.get("product_key"),
        "name": row.get("name"),
        "description": row.get("description"),
        "image_url": row.get("image_url"),
        "url": row.get("product_url"),
        "ctas": ctas,
        "options": options,
    }
