# -*- coding: utf-8 -*- """RAG 知识库存储与检索(Issue #46 / PRD 5.4)。 `RagKnowledgeBase` 是模板 RAG 库的运行时形态:由模板配置 + 文档加载器 自动构建(`from_template_config`),换行业只换模板资产(YAML + 文档集), 内核零改动。检索返回**命中的文档片段 + 来源串**(`RetrievalHit.source`), 满足 PRD 5.4「RAG 答案强制引用溯源」。 检索实现为零依赖的倒排词频匹配(中英混合分词 + 子串/单词计数打分): - 向量化/embedding 由部署侧向量库(如 Milvus)接入,本模块保证 检索语义(召回 + 溯源 + 类别过滤)与向量库一致; - 支持按知识源类别(工艺规范 / SOP / 国标)过滤检索范围。 """ from __future__ import annotations import re from dataclasses import dataclass from typing import Callable, Dict, List, Optional, Sequence from .documents import Chunk, KbDocument, chunk_document from .templating import ( KbTemplateConfig, KbTemplateNaming, KnowledgeSourceKind, SOURCE_KINDS, ) # 中英混合分词:英文单词/数字 + 中文连续块 _TOKEN_RE = re.compile(r"[a-zA-Z0-9]+|[\u4e00-\u9fff]+") # 文档加载器:文档标题 → 原始文本(由部署侧提供:读对象存储/本地目录) DocumentLoader = Callable[[str], str] def _tokenize(text: str) -> List[str]: """分词:英文按单词;中文连续块切**二元组(bigram)**。 bigram 使整句中文查询(如「炉温骤升怎么处置」)与段落中的连续子串 可匹配(共享 bigram 计数),无需外部分词器,零依赖可复现。 """ tokens: List[str] = [] for t in _TOKEN_RE.findall(text or ""): t = t.lower() if t.isascii() or len(t) < 2: tokens.append(t) else: tokens.extend(t[i : i + 2] for i in range(len(t) - 1)) return tokens def _count_token(text_lower: str, token: str) -> int: """chunk 内 token 出现次数:英文按单词边界、中文按子串(查询词原样匹配)。""" if token.isascii(): return len(re.findall(rf"\b{re.escape(token)}\b", text_lower)) return text_lower.count(token) @dataclass class RetrievalHit: """一次检索命中:段落文本 + 来源(引用溯源,PRD 5.4)。""" chunk: Chunk score: float @property def text(self) -> str: return self.chunk.text @property def source(self) -> str: """溯源串(如 `沸腾氯化炉异常处置SOP §2.1 ¶3`),随答案返回给用户。""" return self.chunk.source @property def category(self) -> str: return self.chunk.category def to_dict(self) -> Dict[str, object]: return { "source": self.source, "category": self.category, "text": self.chunk.text, "score": round(self.score, 4), } class RagKnowledgeBase: """模板化 RAG 知识库(内存实现,零外部依赖)。 用法: ```python kb = RagKnowledgeBase.from_template_config( config, loader=lambda title: read_object(title)) hits = kb.search("炉温骤升怎么处理", top_k=3) for h in hits: print(h.source, h.text) # 溯源 + 片段 ``` """ def __init__(self, naming: Optional[KbTemplateNaming] = None): self.naming = naming or KbTemplateNaming("default") self._docs: Dict[str, KbDocument] = {} self._chunks: List[Chunk] = [] # ------------------------------------------------------------------ # 构建 # ------------------------------------------------------------------ def add_document(self, doc: KbDocument, max_chars: int = 500) -> int: """入库一份文档,返回新增段落数。""" if doc.doc_id in self._docs: raise ValueError(f"文档 {doc.doc_id!r} 已存在(同一知识库内 doc_id 唯一)") self._docs[doc.doc_id] = doc chunks = chunk_document(doc, max_chars=max_chars) self._chunks.extend(chunks) return len(chunks) def add_documents(self, docs: Sequence[KbDocument], max_chars: int = 500) -> int: return sum(self.add_document(d, max_chars=max_chars) for d in docs) @classmethod def from_template_config( cls, config: KbTemplateConfig, loader: DocumentLoader, max_chars: int = 500, ) -> "RagKnowledgeBase": """按模板配置构建知识库:遍历三类知识源文档清单,经 loader 取文本入库。 换行业只改模板资产(kb.template.yaml + 文档集),内核零改动。 """ kb = cls(naming=KbTemplateNaming(config.template)) for source in config.sources: for title in source.documents: text = loader(title) if not text or not text.strip(): raise ValueError( f"文档 {title!r}({source.kind.value})加载为空,无法入库" ) doc = KbDocument( doc_id=source.kind.value + ":" + _slug(title), title=title, text=text, category=source.kind.value, version=config.version, ) kb.add_document(doc, max_chars=max_chars) return kb # ------------------------------------------------------------------ # 检索 # ------------------------------------------------------------------ def search( self, query: str, top_k: int = 5, categories: Optional[Sequence[KnowledgeSourceKind]] = None, ) -> List[RetrievalHit]: """检索:query → 命中文档片段(带来源),按相关度降序取 top_k。 `categories` 限定检索范围(如只看 SOP);缺省检索全部知识源。 """ tokens = _tokenize(query) if not tokens or top_k <= 0: return [] cat_set = {c.value for c in categories} if categories else None scored: List[RetrievalHit] = [] for chunk in self._chunks: if cat_set is not None and chunk.category not in cat_set: continue text_lower = chunk.text.lower() freq = sum(_count_token(text_lower, t) for t in tokens) if freq > 0: # TF 密度打分:频次 / 段落长度平方根,抑制长段落重复命中偏好 score = freq / (len(chunk.text) ** 0.5) scored.append(RetrievalHit(chunk=chunk, score=float(score))) scored.sort(key=lambda h: (-h.score, h.chunk.seq)) return scored[:top_k] def category_stats(self) -> Dict[str, int]: """按知识源类别统计段落数(模板配置台展示用)。""" stats: Dict[str, int] = {} for chunk in self._chunks: stats[chunk.category] = stats.get(chunk.category, 0) + 1 return stats # ------------------------------------------------------------------ @property def doc_count(self) -> int: return len(self._docs) @property def chunk_count(self) -> int: return len(self._chunks) @property def template(self) -> str: return self.naming.template def _slug(title: str) -> str: """文档标题 → 对象键安全 slug(小写 + 非安全字符替换为 `_`)。""" from .templating import sanitize return sanitize(title)