#!/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()))