"""03_pgvector_rag.py —— 第 18 章配套代码 #3

用途
    用 pgvector 做一个最小可运行的语义检索 demo（RAG 的「R」部分）：
        1. 用 sentence-transformers/all-MiniLM-L6-v2 把若干文档转成 384 维向量
        2. 写入 ch18_docs.embedding（vector 类型）
        3. 建 HNSW 索引
        4. 输入查询文本 → embedding → ORDER BY embedding <=> query_emb LIMIT 5
    输出最相似的 Top 5 文档及 cosine 距离。

依赖
    1. PG 端：pgvector 扩展（CREATE EXTENSION vector）+ 已建好 ch18_docs 表（init.sql 完成）
    2. Python 端：
           pip install "psycopg[binary]>=3.1" sentence-transformers numpy
       第一次运行会自动从 HuggingFace 下载模型（约 80MB）。

降级方案
    如果机器无法访问 HuggingFace，脚本会回退到「随机 + 关键词哈希」的伪 embedding，
    主要为了演示 SQL 流程，并不能体现真实的语义检索效果。
"""

from __future__ import annotations

import sys
import textwrap

try:
    import psycopg
    from psycopg.rows import tuple_row
except ImportError:
    sys.exit('请先安装 psycopg v3：pip install "psycopg[binary]>=3.1"')


CONN_INFO = "host=127.0.0.1 port=5432 dbname=learn_pg user=postgres"
DIM = 384

DOCS = [
    ("PostgreSQL 入门",
     "PostgreSQL 是一款开源的对象-关系数据库管理系统，强调标准 SQL 与可扩展性。"),
    ("MVCC 多版本并发控制",
     "PG 用元组的 xmin/xmax 字段实现多版本并发控制，读不阻塞写、写不阻塞读。"),
    ("WAL 与崩溃恢复",
     "事务先写 Write-Ahead Log 再改数据页，宕机后通过重放 WAL 把数据库带回一致状态。"),
    ("VACUUM 与死元组",
     "更新和删除会留下死元组，autovacuum 后台进程负责清理并防止 XID 回卷。"),
    ("索引种类",
     "B-Tree 适合等值与范围；GIN 适合 JSONB / 全文 / 数组；GiST 适合空间数据；BRIN 适合时序大表。"),
    ("PostGIS 地理信息",
     "PostGIS 是 PG 的空间扩展，提供 geometry/geography 类型，支持「附近商家」等空间查询。"),
    ("pgvector 向量检索",
     "pgvector 扩展提供 vector 类型，配合 HNSW 索引可在 PG 内做语义检索，是 RAG 的常见后端。"),
    ("Streaming Replication",
     "PG 通过流复制把主库 WAL 实时推送到备库，支持同步、异步与 quorum 模式。"),
    ("Logical Replication",
     "逻辑复制基于 WAL 解码，按表/行级别复制变更，可跨大版本、跨架构。"),
    ("PgBouncer",
     "PgBouncer 是 PG 的轻量级连接池，transaction pooling 能把上万客户端复用到几十个真实连接。"),
    ("Citus 分布式 PG",
     "Citus 让 PG 支持水平分片，把大表按 shard key 切到多个 worker 节点。"),
    ("电商订单设计",
     "订单系统通常包含 users / orders / order_items / payments 几张表，利用 RETURNING 与 ON CONFLICT 简化写入。"),
]


def encode_with_sbert(texts: list[str]):
    """优先用 sentence-transformers 生成真实 embedding。"""
    try:
        from sentence_transformers import SentenceTransformer
        model = SentenceTransformer("sentence-transformers/all-MiniLM-L6-v2")
        embs = model.encode(texts, normalize_embeddings=True)
        return embs.tolist()
    except Exception as exc:
        print(f"⚠️  sentence-transformers 不可用 ({exc.__class__.__name__})，"
              "fallback 到伪 embedding（仅演示 SQL 流程，不代表真实语义）")
        return None


def encode_fake(texts: list[str]):
    """兜底：把字符 hash 映射到 384 维归一化向量。"""
    import hashlib
    import math
    out = []
    for t in texts:
        vec = [0.0] * DIM
        for token in t.lower().replace('，',' ').replace('。',' ').split():
            h = int(hashlib.md5(token.encode('utf-8')).hexdigest(), 16)
            for i in range(8):
                idx = (h >> (i * 4)) & (DIM - 1)
                vec[idx] += 1.0
        norm = math.sqrt(sum(v*v for v in vec)) or 1.0
        out.append([v / norm for v in vec])
    return out


def to_pg_vector(vec) -> str:
    """把 list[float] 转 pgvector 文本字面量 [0.1,0.2,...]"""
    return "[" + ",".join(f"{v:.6f}" for v in vec) + "]"


def main() -> None:
    try:
        conn = psycopg.connect(CONN_INFO, autocommit=True)
    except psycopg.OperationalError as exc:
        sys.exit(f"连接失败：{exc}")

    with conn, conn.cursor(row_factory=tuple_row) as cur:
        cur.execute("SELECT 1 FROM pg_extension WHERE extname = 'vector'")
        if cur.fetchone() is None:
            sys.exit(
                "❌ pgvector 未启用。安装：参见 https://github.com/pgvector/pgvector\n"
                "   通常：\n"
                "     git clone --branch v0.7.4 https://github.com/pgvector/pgvector\n"
                "     cd pgvector && make && sudo make install\n"
                "   然后：CREATE EXTENSION vector;"
            )

        # 1. 编码
        texts = [t + "  " + d for t, d in DOCS]
        embs = encode_with_sbert(texts)
        if embs is None:
            embs = encode_fake(texts)
        else:
            assert len(embs[0]) == DIM, f"模型输出维度 {len(embs[0])} ≠ 表定义 {DIM}"

        # 2. 写入 ch18_docs（先清表）
        cur.execute("TRUNCATE ch18_docs RESTART IDENTITY")
        for (title, content), vec in zip(DOCS, embs):
            cur.execute(
                "INSERT INTO ch18_docs (title, content, embedding) VALUES (%s, %s, %s::vector)",
                (title, content, to_pg_vector(vec)),
            )
        print(f"✅ 已写入 {len(DOCS)} 条文档")

        # 3. 建 HNSW 索引（cosine 距离）
        cur.execute("DROP INDEX IF EXISTS ch18_idx_docs_embedding_hnsw")
        try:
            cur.execute(
                "CREATE INDEX ch18_idx_docs_embedding_hnsw "
                "ON ch18_docs USING hnsw (embedding vector_cosine_ops)"
            )
            print("✅ HNSW 索引已建（vector_cosine_ops）")
        except psycopg.errors.FeatureNotSupported:
            print("⚠️  当前 pgvector 版本不支持 HNSW，回退到 IVFFlat")
            cur.execute(
                "CREATE INDEX ch18_idx_docs_embedding_ivf "
                "ON ch18_docs USING ivfflat (embedding vector_cosine_ops) WITH (lists = 10)"
            )

        # 4. 检索
        queries = [
            "怎么处理 PG 里的死元组？",
            "如何在数据库里做语义搜索？",
            "PG 主从同步原理",
        ]
        for q in queries:
            qemb = encode_with_sbert([q]) or encode_fake([q])
            qvec = to_pg_vector(qemb[0])
            print(f"\n🔍 查询：{q}")
            cur.execute(
                """
                SELECT id, title, content,
                       embedding <=> %s::vector AS cos_dist
                FROM ch18_docs
                ORDER BY embedding <=> %s::vector
                LIMIT 3
                """,
                (qvec, qvec),
            )
            for i, (did, title, content, dist) in enumerate(cur.fetchall(), 1):
                print(f"  {i}) [{dist:.3f}] {title}")
                print(f"       {textwrap.shorten(content, width=70)}")

        # 5. 操作符速查
        print("\n💡 pgvector 距离操作符：")
        print("   · v1 <-> v2  → L2 欧氏距离")
        print("   · v1 <#> v2  → 负内积（值越小越相似）")
        print("   · v1 <=> v2  → cosine 距离（推荐用于文本 embedding）")
        print("   · 索引：HNSW（PG 0.5.0+，召回好且快） / IVFFlat（节省内存，要先 ANALYZE）")


if __name__ == "__main__":
    main()
