Files

207 lines
7.4 KiB
Python
Raw Permalink Normal View History

# -*- 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)