Files
server-core/app/park/rag.py
T

461 lines
18 KiB
Python
Raw Normal View History

# -*- coding: utf-8 -*-
"""园区知识库 —— 轻量本地 RAGembedding + numpy 余弦检索,双路索引)
升级为「细粒度分块 + LLM 生成问答对 + 双路向量索引(原句 + 问题)+ 余弦检索返回原文」:
- 知识源:backend/knowledge/park.md(按 H2 节 → 原子事实两级分块)
- 向量化:dashscope text-embedding-v4(云端 API,索引构建一次后缓存,构建时 L2 归一化)
- 问答:复用 llm._chat_once 对每块生成 3-5 个口语化问句,缓存到 kb_qa.json(无 key 自动单路降级)
- 检索:query 同时比对原句向量与问题向量,按块去重,返回 top_k 个原文块
- 提供:
- retrieve(query, top_k) → 检索命中的原文块(对话动态注入)
- full_kb() → 整库原文(小库全量注入,供语音 instructions
- brief() → 前 N 字摘要(兼容 /api/kb/brief 与检索无命中回退)
"""
import hashlib
import json
import logging
import os
import re
from pathlib import Path
import numpy as np
log = logging.getLogger("dpm.rag")
def _trunc(s, n=300):
"""日志用截断:超长内容只留前 n 字并标注总长。"""
s = str(s)
return s if len(s) <= n else s[:n] + f"...(共{len(s)}字)"
KB_DIR = Path(__file__).resolve().parent / "knowledge"
KB_GLOB = "*.md" # 知识库 = knowledge/ 下全部 markdownpark.md + opc.md 等)
META_FILE = KB_DIR / "kb_meta.json" # 元数据(JSON,可读):model/hash/chunks/questions/embed_texts/embed_owner
VECTOR_FILE = KB_DIR / "kb_vectors.npy" # 向量矩阵(二进制 numpy float32,快载小体积)
LEGACY_INDEX_FILE = KB_DIR / "kb_index.json" # 旧版单文件索引(迁移后删除)
QA_FILE = KB_DIR / "kb_qa.json"
EMBED_MODEL = "text-embedding-v4"
EMBED_BATCH = 10 # DashScope embedding 单批上限
QA_BATCH = 5 # 每批几块生成问答
CHUNK_MAX = 240 # 单块超过此长度按句拆
_index = None # 缓存 {chunks, questions, embed_texts, embed_owner, embeddings(归一化 np.ndarray)}
def _kb_text() -> str:
"""拼接 knowledge/ 下全部 markdown 内容(按文件名排序,稳定顺序)。"""
parts = []
for f in sorted(KB_DIR.glob(KB_GLOB)):
try:
parts.append(f.read_text("utf-8"))
except OSError as e: # noqa: BLE001
log.warning("rag: 跳过知识文件 %s: %s", f.name, e)
return "\n\n".join(parts)
def _data_hash() -> str:
"""全部知识文件内容签名,用于索引/问答缓存失效判断。"""
return hashlib.sha1(_kb_text().encode("utf-8")).hexdigest()
def _split_long(text: str, limit: int = CHUNK_MAX) -> list[str]:
"""超长块按句号/分号/顿号等边界拆成多块,保留语义完整。"""
if len(text) <= limit:
return [text]
parts: list[str] = []
cur = ""
for seg in re.split(r"(?<=[。!?;、])", text):
if len(cur) + len(seg) > limit and cur:
parts.append(cur.strip())
cur = ""
cur += seg
if cur.strip():
parts.append(cur.strip())
return [p for p in parts if len(p) > 8]
def _chunk_md(text: str) -> list[str]:
"""两级切分:H2 节标题作上下文前缀 + 节内原子事实。
企业条目(加粗公司名行 `**...|公司名**` + 负责人/简介列表)合并为单个块,
保证「企业名称 + 企业介绍」作为一个整体可被按名召回。
"""
chunks: list[str] = []
section = ""
para: list[str] = []
company: list[str] | None = None # 正在累积的单个企业块
def flush_para():
nonlocal para
if para:
t = "\n".join(para).strip()
if len(t) > 8:
chunks.extend(_split_long(t))
para = []
def flush_company():
nonlocal company
if company:
t = "\n".join(company).strip()
if len(t) > 8:
chunks.append(t) # 企业块不拆,保证名称+介绍同块
company = None
def add_fact(content: str):
full = f"{section}\n{content}" if section else content
if len(full) > 8:
chunks.extend(_split_long(full))
for line in text.splitlines():
s = line.rstrip()
if s.startswith("# "): # H1:不产块
flush_para(); flush_company()
continue
if s.startswith("## "): # H2 小节:切换上下文前缀
flush_para(); flush_company()
section = s[3:].strip()
continue
if s.startswith("###"): # H3 子节:上下文(如分区),flush 企业
flush_para(); flush_company()
para.append(s.strip())
continue
stripped = s.strip()
if not stripped:
continue
# 企业条目:加粗公司名行作为单个企业块的起点
if stripped.startswith("**") and "" in stripped:
flush_para(); flush_company()
company = [stripped]
continue
if company is not None:
# 企业块的跟随行(负责人/简介等列表项)
if stripped.startswith("- ") or stripped.startswith("问:") or re.match(r"^\d+[\.、]\s", stripped):
company.append(stripped)
else:
flush_company()
para.append(stripped)
continue
if stripped.startswith("- "): # 列表项 = 原子事实
flush_para(); add_fact(stripped)
elif re.match(r"^\d+[\.、]\s", stripped): # 编号步
flush_para(); add_fact(stripped)
elif stripped.startswith("问:"): # 单条问答
flush_para(); add_fact(stripped)
else: # 普通段落:累积后整段成块(超长再拆)
para.append(stripped)
flush_para()
flush_company()
return chunks
def _company_name(chunk: str) -> str:
"""从企业块提取企业名称:`**加速1|云南派音人工智能科技有限公司**` → `云南派音人工智能科技有限公司`。"""
m = re.match(r"^\*\*[^\n]*([^*]+)\*\*", chunk) or re.match(r"^\*\*([^*]+)\*\*", chunk)
return m.group(1).strip() if m else ""
def _embed(texts: list[str]) -> list[list[float]]:
"""dashscope text-embedding-v4 批量向量化(每批 ≤10 条,按 text_index 排序拼接)。"""
import dashscope
api_key = os.environ.get("DASHSCOPE_API_KEY") or ""
out: list[list[float]] = []
for i in range(0, len(texts), EMBED_BATCH):
batch = texts[i:i + EMBED_BATCH]
resp = dashscope.TextEmbedding.call(
model=EMBED_MODEL,
input=batch,
api_key=api_key,
)
if resp.status_code != 200:
raise RuntimeError(
f"embedding API {resp.status_code}: {resp.code} {resp.message}"
)
embs = sorted(
resp.output["embeddings"],
key=lambda e: e.get("text_index", e.get("index", 0)),
)
out.extend(e["embedding"] for e in embs)
log.info("rag: embedding 完成 %d 条(%s,每批 %d", len(texts), EMBED_MODEL, EMBED_BATCH)
return out
def _normalize(vectors: list[list[float]]) -> np.ndarray:
"""转为 float32 np.ndarray 并做 L2 行归一化(点积即余弦相似度)。"""
arr = np.asarray(vectors, dtype=np.float32)
norms = np.linalg.norm(arr, axis=1, keepdims=True)
norms[norms == 0] = 1.0
return arr / norms
# ── LLM 问答生成 ────────────────────────────────────────────────
def _parse_json_array(text: str) -> list:
"""解析模型返回的 JSON 数组:剥 markdown 围栏 → json.loads → 截取首尾方括号兜底。"""
text = (text or "").strip()
if text.startswith("```"):
lines = [l for l in text.splitlines() if not l.strip().startswith("```")]
text = "\n".join(lines).strip()
for candidate in (text,):
try:
return json.loads(candidate)
except Exception:
pass
i = text.find("[")
j = text.rfind("]")
if i != -1 and j > i:
try:
return json.loads(text[i:j + 1])
except Exception:
pass
return []
def _qa_for_batch(chunks: list[str]) -> list[list[str]]:
"""对一批分块,用 LLM 生成每块的 3-5 个口语化问句(输出 JSON 数组,与块一一对应)。"""
from .llm import _chat_once
numbered = "\n\n".join(f"[{i}]\n{c}" for i, c in enumerate(chunks))
def build_prompt():
return (
"下面是园区知识库的若干分块,每块标有 [编号]。"
"请为每一块生成 3-5 个用户可能会问的自然口语化中文问句,"
"覆盖该块全部信息点与常见同义问法。"
f"严格只输出一个 JSON 数组,数组长度必须为 {len(chunks)}"
"第 i 个元素是第 i 块的问题字符串数组。不要输出任何其他文字或解释。\n\n"
+ numbered
)
for attempt in range(2):
try:
data = _chat_once([{"role": "user", "content": build_prompt()}], with_tools=False)
content = data["choices"][0]["message"].get("content") or ""
arr = _parse_json_array(content)
if isinstance(arr, list) and arr:
out: list[list[str]] = []
for i in range(len(chunks)):
qs = arr[i] if i < len(arr) and isinstance(arr[i], list) else []
out.append([str(q).strip() for q in qs if str(q).strip()][:5])
return out
except Exception as e: # noqa: BLE001
log.warning("rag: QA 批生成失败(尝试 %d/%d: %s", attempt + 1, 2, e)
return [[] for _ in chunks]
def _first_line(chunk: str) -> str:
"""取块首行(去标题),用作无问句时的兜底检索问题。"""
for ln in chunk.splitlines():
ln = ln.strip()
if ln and not ln.startswith("#"):
return ln
return chunk[:60]
def _ensure_questions(chunks: list[str], questions) -> list[list[str]]:
"""确保每块都有至少一个检索问题:未生成问句的块,用块首句兜底(保证全量覆盖)。"""
qs: list[list[str]] = [list(q or []) for q in (questions or [])]
if len(qs) < len(chunks):
qs.extend([] for _ in range(len(chunks) - len(qs)))
for i, c in enumerate(chunks):
if not qs[i]:
qs[i] = [_first_line(c)]
return qs
def _load_questions(chunks: list[str]) -> list[list[str]]:
"""加载/生成问答缓存;返回与 chunks 对齐的问题数组(保证每块 ≥1 个问题)。无 key 时单路回退。"""
from .config import settings
if not settings.DASHSCOPE_API_KEY:
log.warning("rag: 无 DASHSCOPE_API_KEY,跳过问答生成(单路检索)")
return _ensure_questions(chunks, [[] for _ in chunks])
h = _data_hash()
if QA_FILE.exists():
try:
data = json.loads(QA_FILE.read_text("utf-8"))
if data.get("data_hash") == h and data.get("chunks") == chunks:
log.info("rag: 问答缓存已加载(%d 块)", len(chunks))
return _ensure_questions(chunks, data["questions"])
except Exception: # noqa: BLE001
log.warning("rag: 问答缓存损坏,重建")
questions: list[list[str]] = []
total = len(chunks)
for i in range(0, total, QA_BATCH):
questions.extend(_qa_for_batch(chunks[i:i + QA_BATCH]))
log.info("rag: 问答生成进度 %d/%d", min(i + QA_BATCH, total), total)
questions = _ensure_questions(chunks, questions)
QA_FILE.write_text(
json.dumps({"data_hash": h, "chunks": chunks, "questions": questions}, ensure_ascii=False),
"utf-8",
)
log.info("rag: 问答缓存已生成(%d 块)-> %s", len(chunks), QA_FILE)
return questions
# ── 双路向量索引 ────────────────────────────────────────────────
def _migrate_legacy():
"""把旧版单文件 kb_index.json(向量内嵌在 JSON)就地转为 kb_vectors.npy + kb_meta.json,再删除旧文件。"""
if not LEGACY_INDEX_FILE.exists() or META_FILE.exists():
return
try:
d = json.loads(LEGACY_INDEX_FILE.read_text("utf-8"))
np.save(VECTOR_FILE, np.asarray(d["embeddings"], dtype=np.float32))
META_FILE.write_text(
json.dumps({
"model": d.get("model"),
"data_hash": d.get("data_hash"),
"chunks": d["chunks"],
"questions": d.get("questions", []),
"embed_texts": d["embed_texts"],
"embed_owner": d["embed_owner"],
}, ensure_ascii=False),
"utf-8",
)
LEGACY_INDEX_FILE.unlink()
log.info("rag: 已迁移旧 kb_index.json -> kb_vectors.npy + kb_meta.json(无重建,未调 API")
except Exception as e: # noqa: BLE001
log.warning("rag: 旧索引迁移失败(将重建): %s", e)
def _load_index():
"""构建/加载双路向量索引(原句 + 生成问题)。向量存二进制 .npy,元数据存 JSON。"""
global _index
if _index is not None:
return _index
_migrate_legacy()
h = _data_hash()
if META_FILE.exists() and VECTOR_FILE.exists():
try:
data = json.loads(META_FILE.read_text("utf-8"))
valid = (
data.get("model") == EMBED_MODEL
and data.get("data_hash") == h
and len(data.get("embed_texts", [])) == len(data.get("embed_owner", []))
)
if valid:
vectors = np.load(VECTOR_FILE) # 已归一化 float32
if len(vectors) == len(data["embed_texts"]):
_index = {
"chunks": data["chunks"],
"questions": data.get("questions", []),
"embed_texts": data["embed_texts"],
"embed_owner": data["embed_owner"],
"embeddings": vectors,
}
log.info("rag: 双路知识索引已加载(%d 块 → %d 向量, %s",
len(_index["chunks"]), len(_index["embed_texts"]), EMBED_MODEL)
return _index
log.warning("rag: 索引失效(model/hash/长度不符),重建")
except Exception: # noqa: BLE001
log.warning("rag: 索引损坏,重建")
# 构建
chunks = _chunk_md(_kb_text())
questions = _load_questions(chunks)
embed_texts: list[str] = []
embed_owner: list[int] = []
for i, c in enumerate(chunks):
embed_texts.append(c)
embed_owner.append(i)
# 企业块:把「企业名称」作为确定性问题加入索引,保证按名称可召回(名称为问、介绍为答)
cname = _company_name(c)
if cname:
embed_texts.append(cname)
embed_owner.append(i)
for q in questions[i]:
embed_texts.append(q)
embed_owner.append(i)
log.info("rag: 构建双路索引(%d 块 → %d 向量,%s embedding 中…)",
len(chunks), len(embed_texts), EMBED_MODEL)
arr = _normalize(_embed(embed_texts))
_index = {
"chunks": chunks,
"questions": questions,
"embed_texts": embed_texts,
"embed_owner": embed_owner,
"embeddings": arr,
}
np.save(VECTOR_FILE, arr)
META_FILE.write_text(
json.dumps({
"model": EMBED_MODEL,
"data_hash": h,
"chunks": chunks,
"questions": questions,
"embed_texts": embed_texts,
"embed_owner": embed_owner,
}, ensure_ascii=False),
"utf-8",
)
log.info("rag: 双路索引构建完成(%d 块, dim=%d-> %s + %s",
len(chunks), arr.shape[1], VECTOR_FILE, META_FILE)
return _index
# ── 对外接口 ────────────────────────────────────────────────────
def retrieve(query: str, top_k: int = 3, threshold: float = 0.3) -> list[str]:
"""双路检索:query 同时比对原句与问题向量,按块去重,返回 top_k 个原文块。"""
if not query or not query.strip():
return []
idx = _load_index()
q = _normalize([_embed([query])[0]])[0]
sims = np.dot(idx["embeddings"], q)
order = np.argsort(-sims)
best: dict[int, float] = {} # owner -> 该块最高分
for e in order:
owner = idx["embed_owner"][int(e)]
sc = float(sims[int(e)])
if sc < threshold:
break
if owner not in best:
best[owner] = sc
if len(best) >= top_k:
break
ordered = sorted(best.items(), key=lambda kv: kv[1], reverse=True)
hits = [idx["chunks"][o] for o, _ in ordered]
log.info("rag: retrieve() query=%s top_k=%d threshold=%.2f -> 命中 %d",
_trunc(query, 120), top_k, threshold, len(hits))
for j, (o, sc) in enumerate(ordered):
log.info("rag: [%d] %.4f %s", j, sc, _trunc(idx["chunks"][o], 300))
return hits
def brief(max_chars: int = 1200) -> str:
"""固定知识摘要(前 N 字;兼容 /api/kb/brief 与检索无命中回退)。"""
idx = _load_index()
out = []
total = 0
for c in idx["chunks"]:
if total + len(c) > max_chars:
break
out.append(c)
total += len(c)
text = "\n\n".join(out)
log.info("rag: brief() 知识库摘要 %d 字 -> %s", total, _trunc(text, 300))
return text
def full_kb() -> str:
"""整库原文(小库全量注入,用于语音 instructions,修复尾部知识丢失)。"""
idx = _load_index()
text = "\n\n".join(idx["chunks"])
log.info("rag: full_kb() 全库 %d 字 / %d", len(text), len(idx["chunks"]))
return text
def build_instructions() -> str:
"""语音助手的完整基础提示词 = 配置提示词 + 整库知识(服务端统一组装)。"""
from .config import settings
base = settings.S2S_INSTRUCTIONS.strip()
try:
kb = full_kb()
if kb:
base += f"\n\n【园区知识库】\n{kb}"
except Exception as e: # noqa: BLE001
log.warning("rag: 知识库拼接失败(仅用基础提示词): %s", e)
log.info("rag: build_instructions() 语音系统提示词=%s", _trunc(base, 600))
return base