"""语义检索：pgvector + HNSW 索引

流程：
    1. 用 sentence-transformers 把文章正文 embed 成 384 维向量
    2. 写入 post_embeddings 表（vector(384) 列 + HNSW 索引）
    3. 用 `<=>`（余弦距离）做近邻检索

依赖：sentence-transformers, numpy  （可选；没有时会 fallback 到随机向量 demo）
运行：python code/search_vector.py "如何做向量检索？"
"""
from __future__ import annotations

import os
import sys
from typing import List

from db import fetch_all, fetch_one, execute, close_pool

EMBED_DIM = 384
MODEL_NAME = "all-MiniLM-L6-v2"


# ------------------------------- embedding --------------------------------
def _get_encoder():
    """懒加载 embedding 模型。没装 sentence-transformers 就返回 None。"""
    try:
        from sentence_transformers import SentenceTransformer
    except ImportError:
        return None
    global _encoder
    if "_encoder" not in globals() or _encoder is None:
        print(f"[init] 加载 {MODEL_NAME} ...")
        _encoder = SentenceTransformer(MODEL_NAME)
    return _encoder


def embed(text: str) -> List[float]:
    enc = _get_encoder()
    if enc is None:
        # fallback：随机向量（仅供流程演示，不可用于真实检索）
        import random
        random.seed(hash(text) & 0xFFFFFFFF)
        return [random.random() for _ in range(EMBED_DIM)]
    vec = enc.encode(text, normalize_embeddings=True)
    return vec.tolist()


def _to_pg_vector(vec: List[float]) -> str:
    """pgvector 文本表示：'[0.1,0.2,...]'"""
    return "[" + ",".join(f"{x:.6f}" for x in vec) + "]"


# -------------------------------- 入库 ------------------------------------
def upsert_embedding(post_id: int, text: str) -> None:
    vec = embed(text)
    execute(
        """INSERT INTO post_embeddings(post_id, model, embedding)
           VALUES (%s, %s, %s::vector)
           ON CONFLICT (post_id, model) DO UPDATE
              SET embedding = EXCLUDED.embedding,
                  updated_at = now()""",
        (post_id, MODEL_NAME, _to_pg_vector(vec)),
    )


def build_index_for_all(tenant_id: int) -> int:
    """为该租户所有已发布文章生成 embedding。"""
    rows = fetch_all(
        """SELECT id, title, body
           FROM   posts
           WHERE  tenant_id = %s AND status = 'published'
             AND  NOT EXISTS (
                    SELECT 1 FROM post_embeddings e
                    WHERE  e.post_id = posts.id AND e.model = %s)""",
        (tenant_id, MODEL_NAME),
    )
    for r in rows:
        upsert_embedding(r["id"], f"{r['title']}\n{r['body']}")
    return len(rows)


# -------------------------------- 检索 ------------------------------------
def search(tenant_id: int, query: str, limit: int = 5) -> list[dict]:
    """用 `<=>` 余弦距离做最近邻。HNSW 索引会命中 idx_post_emb_hnsw。"""
    qvec = _to_pg_vector(embed(query))
    return fetch_all(
        """SELECT p.id, p.title,
                  (e.embedding <=> %s::vector) AS distance
           FROM   post_embeddings e
           JOIN   posts p ON p.id = e.post_id
           WHERE  p.tenant_id = %s AND e.model = %s
           ORDER  BY e.embedding <=> %s::vector
           LIMIT  %s""",
        (qvec, tenant_id, MODEL_NAME, qvec, limit),
    )


def explain(tenant_id: int, query: str) -> list[dict]:
    qvec = _to_pg_vector(embed(query))
    return fetch_all(
        """EXPLAIN (ANALYZE, BUFFERS, FORMAT TEXT)
           SELECT post_id
           FROM   post_embeddings e
           ORDER  BY embedding <=> %s::vector
           LIMIT  5""",
        (qvec,),
    )


# -------------------------------- demo ------------------------------------
def _demo(query: str) -> None:
    enc = _get_encoder()
    if enc is None:
        print("[WARN] 未安装 sentence-transformers，仅做流程演示（结果不准）。")
        print("       可运行：pip install sentence-transformers")

    n = build_index_for_all(1)
    print(f"[OK] 新生成 {n} 条 embedding（已入库的会跳过）")

    print(f"\n=== 语义检索：{query!r} ===")
    for r in search(1, query, limit=5):
        print(f" dist={r['distance']:.4f}  #{r['id']:>3}  {r['title'][:50]}")

    print("\n-- EXPLAIN（验证 HNSW 索引是否生效）--")
    for r in explain(1, query):
        print(" ", r["QUERY PLAN"])


if __name__ == "__main__":
    q = sys.argv[1] if len(sys.argv) > 1 else "如何实现向量检索?"
    try:
        _demo(q)
    finally:
        close_pool()
