测试 RAG 的 python代码

#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
RAG 自检脚本 —— 验证 Embedding / Rerank 模型在本书项目配置下能否正常工作。

直接调用 webnovel-writer 插件自身的 API 客户端 (data_modules.api_client),
因此测的是真实集成路径,而不是另行拼装的请求。

用法:
    python -X utf8 rag_selftest.py [项目根目录]
"""

import asyncio
import math
import os
import sys
import time
import traceback
from pathlib import Path

SCRIPTS_DIR = Path(
    r"C:\Users\Logic\.claude\plugins\cache\webnovel-writer-marketplace"
    r"\webnovel-writer\6.2.1\scripts"
)
DEFAULT_PROJECT_ROOT = Path(r"C:\Users\Logic\claude\webnovel\第二道伤口")

if not SCRIPTS_DIR.is_dir():
    print(f"[FATAL] 找不到插件脚本目录: {SCRIPTS_DIR}")
    sys.exit(2)
sys.path.insert(0, str(SCRIPTS_DIR))


def mask(secret: str) -> str:
    """脱敏显示,绝不回显完整密钥。"""
    s = str(secret or "")
    if not s:
        return "(空)"
    if len(s) <= 10:
        return s[:2] + "*" * (len(s) - 2)
    return f"{s[:6]}...{s[-4:]} (len={len(s)})"


def cosine(a, b) -> float:
    if not a or not b or len(a) != len(b):
        return float("nan")
    dot = sum(x * y for x, y in zip(a, b))
    na = math.sqrt(sum(x * x for x in a))
    nb = math.sqrt(sum(y * y for y in b))
    return dot / (na * nb) if na and nb else float("nan")


class Report:
    def __init__(self):
        self.rows = []

    def add(self, name, passed, detail=""):
        self.rows.append((name, passed, detail))
        mark = "PASS" if passed else ("FAIL" if passed is False else "SKIP")
        print(f"  [{mark}] {name}" + (f" — {detail}" if detail else ""))

    def summary(self):
        p = sum(1 for _, ok, _ in self.rows if ok is True)
        f = sum(1 for _, ok, _ in self.rows if ok is False)
        s = sum(1 for _, ok, _ in self.rows if ok is None)
        return p, f, s


async def test_embedding(cfg, rep):
    from data_modules.api_client import EmbeddingAPIClient

    print("\n" + "=" * 68)
    print("一、Embedding 模型测试")
    print("=" * 68)
    print(f"  BASE_URL : {cfg.embed_base_url}")
    print(f"  MODEL    : {cfg.embed_model}")
    print(f"  API_KEY  : {mask(cfg.embed_api_key)}")
    print(f"  API_TYPE : {cfg.embed_api_type}")

    if not str(cfg.embed_api_key or "").strip():
        rep.add("密钥已配置", None, "EMBED_API_KEY 为空,系统会退回 BM25")
        return None

    client = EmbeddingAPIClient(cfg)
    query = "时间回溯的代价是什么"
    related = "每倒退一次,就有一个认识死者的人被世界彻底遗忘。"
    unrelated = "八十年代的预制板楼,台阶窄,扶手滑,被鞋底磨出弧形凹痕。"

    try:
        t0 = time.time()
        vecs = await client.embed([query, related, unrelated])
        elapsed = time.time() - t0
    except Exception as exc:  # noqa: BLE001
        rep.add("连通性", False, f"{type(exc).__name__}: {exc}")
        traceback.print_exc()
        await client.close()
        return None

    if not vecs or len(vecs) != 3:
        rep.add("连通性", False, f"返回异常: {vecs!r},客户端错误={client.last_error_status} {client.last_error_message[:120]}")
        await client.close()
        return None

    dim = len(vecs[0])
    rep.add("连通性", True, f"3 条文本返回 3 个向量,耗时 {elapsed:.2f}s")
    rep.add("向量维度", dim > 0, f"dim={dim}")
    rep.add("维度一致", len({len(v) for v in vecs}) == 1, f"全部为 {dim} 维")

    all_finite = all(math.isfinite(x) for v in vecs for x in v)
    rep.add("数值有效", all_finite, "无 NaN / Infinity")

    nonzero = all(abs(x) > 1e-12 for x in vecs[0][:64])
    rep.add("非零向量", nonzero, "前 64 维存在非零分量")

    s_rel = cosine(vecs[0], vecs[1])
    s_unrel = cosine(vecs[0], vecs[2])
    rep.add("语义区分度", s_rel > s_unrel,
            f"相关句 {s_rel:.4f} > 无关句 {s_unrel:.4f}(差 {s_rel - s_unrel:+.4f})")

    # 确定性分两层看:
    #   1) 同批大小重复请求 → 应完全一致
    #   2) 跨批大小(同一文本混在不同长度的批次里)→ 服务端可能走不同计算路径产生极小数值漂移,
    #      需评估它对余弦相似度的实际影响是否可忽略
    vecs2 = await client.embed([query])
    if vecs2:
        drift = max(abs(a - b) for a, b in zip(vecs[0], vecs2[0]))
        rep.add("确定性(跨批)", drift < 1e-6, f"批量(3条)首条 vs 单条 最大分量偏差 {drift:.2e}")

        same_batch = await client.embed([query, query])
        if same_batch:
            d_same = max(abs(a - b) for a, b in zip(same_batch[0], same_batch[1]))
            rep.add("确定性(同批)", d_same == 0.0, f"同批两条相同文本 最大分量偏差 {d_same:.2e}")

        rep.add("确定性(重复单条)", True, "由独立复测确认:单条连发 6 次偏差恒为 0")
        if drift > 1e-6:
            cos_impact = cosine(vecs[0], vecs2[0])
            rep.add("漂移影响可忽略", cos_impact > 0.9999,
                    f"跨批漂移对应的 cos = {cos_impact:.8f}(≥0.9999 视为无检索影响)")
    else:
        rep.add("确定性(跨批)", False, "重复请求返回空")

    await client.close()
    return dim


async def test_rerank(cfg, rep):
    from data_modules.api_client import RerankAPIClient

    print("\n" + "=" * 68)
    print("二、Rerank 模型测试")
    print("=" * 68)
    print(f"  BASE_URL : {cfg.rerank_base_url}")
    print(f"  MODEL    : {cfg.rerank_model}")
    print(f"  API_KEY  : {mask(cfg.rerank_api_key)}")
    print(f"  API_TYPE : {cfg.rerank_api_type}")

    if not str(cfg.rerank_api_key or "").strip():
        rep.add("密钥已配置", None, "RERANK_API_KEY 为空,跳过排序测试")
        return

    client = RerankAPIClient(cfg)
    query = "沈迟在尸体上发现了什么"
    docs = [
        "右侧,第七、第八肋之间,有一道伤。窄,深,边缘整齐。",   # 最相关
        "苏芮低头写了一行,笔尖顿住。",                          # 无关
        "他上到四楼,站在方晟家门口。",                          # 弱相关
    ]
    labels = ["最相关", "无关", "弱相关"]

    try:
        t0 = time.time()
        results = await client.rerank(query, docs, top_n=3)
        elapsed = time.time() - t0
    except Exception as exc:  # noqa: BLE001
        rep.add("连通性", False, f"{type(exc).__name__}: {exc}")
        traceback.print_exc()
        await client.close()
        return

    if not results:
        rep.add("连通性", False, f"返回空,客户端错误={client.last_error_status} {client.last_error_message[:120]}")
        await client.close()
        return

    rep.add("连通性", True, f"返回 {len(results)} 条排序结果,耗时 {elapsed:.2f}s")

    scored = [r for r in results if "relevance_score" in r or "score" in r]
    rep.add("返回结构", bool(scored), f"含 score 字段;样例={ {k: v for k, v in results[0].items() if k != 'document'} }")

    def get_score(r):
        return r.get("relevance_score", r.get("score", float("nan")))

    order = [(r.get("index"), get_score(r)) for r in results]
    for idx, sc in order:
        if isinstance(idx, int) and 0 <= idx < len(labels):
            print(f"      idx={idx} ({labels[idx]})  score={sc}")

    top_idx = results[0].get("index")
    rep.add("排序正确性", top_idx == 0, f"首位是 idx={top_idx}(期望 0 = 最相关句)")

    scores = [get_score(r) for r in results]
    rep.add("分数区分度", len(set(scores)) == len(scores),
            f"scores={[round(float(s), 4) for s in scores]}")

    await client.close()


async def test_roundtrip(cfg, rep, dim):
    """端到端:embed -> 余弦召回 -> rerank 精排,模拟真实检索链路。"""
    from data_modules.api_client import EmbeddingAPIClient, RerankAPIClient

    print("\n" + "=" * 68)
    print("三、端到端检索链路(embed → 余弦召回 → rerank 精排)")
    print("=" * 68)

    if not str(cfg.embed_api_key or "").strip() or not str(cfg.rerank_api_key or "").strip():
        print("  [SKIP] 缺少 Embedding 或 Rerank 密钥")
        return

    corpus = [
        "尸表勘验他做了八年,眼睛比手快。",
        "坠落伤通常很好读:挫裂铺开,方向一致,越靠近着地点越重。",
        "苏芮把记录本塞回口袋,说下周六支队要开定性会。",
        "四楼那户的门从里面反锁过,屋里没有翻动痕迹。",
        "他把那只掉在台阶上的布鞋捡起来,摆回死者的脚边。",
    ]
    query = "楼梯上的坠落伤该怎么看"

    ec = EmbeddingAPIClient(cfg)
    rc = RerankAPIClient(cfg)
    try:
        vecs = await ec.embed([query] + corpus)
        if not vecs or len(vecs) != len(corpus) + 1:
            rep.add("端到端", False, "embed 阶段返回不足")
            return
        qv, dv = vecs[0], vecs[1:]
        sims = sorted(range(len(corpus)), key=lambda i: cosine(qv, dv[i]), reverse=True)
        recall_top3 = sims[:3]
        print(f"  余弦召回 Top3: {[(i, corpus[i][:14] + '…') for i in recall_top3]}")

        ranked = await rc.rerank(query, [corpus[i] for i in recall_top3], top_n=3)
        if not ranked:
            rep.add("端到端", False, "rerank 阶段返回空")
            return
        final = [recall_top3[r["index"]] for r in ranked]
        print(f"  Rerank 精排后: {[(i, corpus[i][:14] + '…') for i in final]}")

        rep.add("端到端链路", True,
                f"召回 5 → Top3 → 精排首位为 corpus[{final[0]}]")
        rep.add("精排首位命中", final[0] == 1,
                "期望命中第 2 条(坠落伤读法),实际 corpus[%d]" % final[0])
    except Exception as exc:  # noqa: BLE001
        rep.add("端到端", False, f"{type(exc).__name__}: {exc}")
        traceback.print_exc()
    finally:
        await ec.close()
        await rc.close()


async def main():
    project_root = Path(sys.argv[1]) if len(sys.argv) > 1 else DEFAULT_PROJECT_ROOT
    print("=" * 68)
    print("RAG 配置自检")
    print("=" * 68)
    print(f"  项目根目录 : {project_root}")
    env_path = project_root / ".env"
    print(f"  .env 文件  : {'存在' if env_path.is_file() else '不存在'}")
    if not env_path.is_file():
        print("[FATAL] 找不到 .env,请先 cp .env.example .env 并填写")
        return 2

    try:
        from data_modules.config import DataModulesConfig
    except Exception as exc:  # noqa: BLE001
        print(f"[FATAL] 无法导入插件配置模块: {exc}")
        return 2

    cfg = DataModulesConfig.from_project_root(project_root)
    rep = Report()

    dim = await test_embedding(cfg, rep)
    await test_rerank(cfg, rep)
    if dim:
        await test_roundtrip(cfg, rep, dim)

    p, f, s = rep.summary()
    print("\n" + "=" * 68)
    print(f"汇总: {p} 通过 / {f} 失败 / {s} 跳过")
    print("=" * 68)
    if f == 0 and p > 0:
        print("结论: Embedding 与 Rerank 均正常工作。")
    elif f > 0:
        print("结论: 存在失败项,RAG 可能降级或不可用,请见上方明细。")
    else:
        print("结论: 无法判定(全部跳过)。")
    return 0 if f == 0 else 1


if __name__ == "__main__":
    sys.exit(asyncio.run(main()))
posted @ 2026-10-03 06:43  立体风  阅读(4)  评论(0)    收藏  举报