Files
rag/knowledge/base.py
lacerate551 100d1a06eb init: RAG 知识库服务初始提交
- 后端 API(Flask + Gunicorn)
- RAG 引擎(混合检索 + 云端 Reranker + 引用溯源)
- 文档解析(MinerU + 多格式支持)
- Docker 生产部署配置
- 排除前端项目、敏感配置、模型文件
2026-06-04 17:35:27 +08:00

470 lines
14 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
知识库管理器 - 基础模块
包含:
- 配置常量
- 数据类定义
- 辅助函数
- BM25Index 类
"""
import os
import json
import pickle
import logging
from typing import List, Dict, Optional, Tuple
from dataclasses import dataclass
from pathlib import Path
import numpy as np
from rank_bm25 import BM25Okapi
import jieba
# 设置日志
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger(__name__)
# ==================== 配置常量 ====================
# 向量存储基础路径(位于 knowledge/vector_store/
VECTOR_STORE_BASE_PATH = os.path.join(
os.path.dirname(os.path.abspath(__file__)),
"vector_store"
)
# 向量库基础路径ChromaDB 数据存储)
CHROMA_DB_BASE_PATH = os.path.join(VECTOR_STORE_BASE_PATH, "chroma")
# BM25 索引基础路径
BM25_INDEX_BASE_PATH = os.path.join(VECTOR_STORE_BASE_PATH, "bm25")
# 向量库元数据文件
KB_METADATA_FILE = "kb_metadata.json"
# 预定义的公开知识库名称
PUBLIC_KB_NAME = "public_kb"
# 默认部门列表
DEFAULT_DEPARTMENTS = ["finance", "hr", "tech", "operation", "marketing"]
# 部门名称映射
DEPARTMENT_NAME_MAP = {
"财务部": "finance", "财务": "finance",
"人事部": "hr", "人事": "hr", "人力资源部": "hr", "人力资源": "hr",
"技术部": "tech", "技术": "tech", "研发部": "tech", "研发": "tech",
"运营部": "operation", "运营": "operation",
"市场部": "marketing", "市场": "marketing",
"法务部": "legal", "法务": "legal",
"行政部": "admin", "行政": "admin",
"finance": "finance", "hr": "hr", "tech": "tech",
"operation": "operation", "marketing": "marketing", "legal": "legal", "admin": "admin",
}
# ==================== 数据结构 ====================
@dataclass
class CollectionInfo:
"""向量库信息"""
name: str
display_name: str
document_count: int = 0
created_at: str = ""
department: str = ""
description: str = ""
@dataclass
class SearchResult:
"""检索结果"""
ids: List[str]
documents: List[str]
metadatas: List[dict]
distances: List[float]
collection_name: str = ""
# ==================== 辅助函数 ====================
def _get_doc_type(filename: str) -> str:
"""
根据文件扩展名判断文档类型
Args:
filename: 文件名,包含扩展名
Returns:
文档类型字符串: 'pdf' | 'word' | 'excel' | 'ppt' | 'other'
Example:
>>> _get_doc_type("report.pdf")
'pdf'
>>> _get_doc_type("data.xlsx")
'excel'
"""
ext = Path(filename).suffix.lower()
type_map = {
'.pdf': 'pdf',
'.docx': 'word', '.doc': 'word',
'.xlsx': 'excel', '.xls': 'excel',
'.pptx': 'ppt', '.ppt': 'ppt',
}
return type_map.get(ext, 'other')
def _extract_figure_number(caption: str, section: str = '') -> str:
"""
从 caption 或 section 中提取图号(增强版)
支持格式:
- 图2.4, 图2-4
- Fig.2.4, Fig 2.4, Figure 2.4
- 图2
Args:
caption: 图片标题/说明文字
section: 章节信息(可选)
Returns:
图号字符串,如 "2.4";未找到返回空字符串
Example:
>>> _extract_figure_number("图2.4 系统架构图")
'2.4'
>>> _extract_figure_number("参见Figure 3.1所示")
'3.1'
"""
import re
text = f"{caption} {section}"
patterns = [
r'\s*(\d+[\.\-]\d+)',
r'Fig\.?\s*(\d+[\.\-]\d+)',
r'Figure\s*(\d+[\.\-]\d+)',
r'[(]\s*图\s*(\d+)\s*[)]',
]
for pattern in patterns:
match = re.search(pattern, text, re.IGNORECASE)
if match:
return match.group(1).replace('-', '.')
return ""
def normalize_department_name(department: str) -> str:
"""
将部门名称标准化为英文标识
支持中文部门名(如"财务部")和英文标识(如"finance"
返回符合 ChromaDB 命名规范的英文标识。
Args:
department: 原始部门名称(中文或英文)
Returns:
标准化的英文标识;无法识别时返回空字符串
Example:
>>> normalize_department_name("财务部")
'finance'
>>> normalize_department_name("tech")
'tech'
"""
if not department:
return ""
if department in DEPARTMENT_NAME_MAP:
return DEPARTMENT_NAME_MAP[department]
if department.replace("_", "").replace("-", "").isalnum() and department.isascii():
return department.lower()
logger.warning(f"无法识别的部门名称: {department}")
return ""
def _extract_section(section_path: str, max_levels: int = 3) -> str:
"""
动态截断章节路径,保留核心+末尾层级
用于在向量检索时提供简洁的章节上下文,
避免过长的章节路径影响语义匹配。
Args:
section_path: 完整章节路径,用 '>' 分隔
max_levels: 最多保留的层级数默认3级
Returns:
截断后的章节路径
Example:
>>> _extract_section("第一章 > 1.1 概述 > 1.1.1 背景 > 1.1.1.1 详细说明")
'1.1 概述 > 1.1.1 背景 > 1.1.1.1 详细说明'
"""
parts = [p.strip() for p in section_path.split('>') if p.strip()]
if len(parts) <= max_levels:
return ' > '.join(parts)
return ' > '.join(parts[-max_levels:])
def _build_semantic_content_for_text(chunk, page_info: dict, doc_type: str) -> str:
"""
构建语义增强内容(文本类型)
将原始文本切片转换为适合向量检索的语义增强格式,
包含标题、章节上下文和正文内容。
PDF 和 Word 差异化处理:
- PDF: 有 text_level、bbox标题识别准确
- Word: 无 text_level、bbox依赖启发式识别
Args:
chunk: 文档切片对象,含 title、content、text_level 属性
page_info: 页面信息字典,含 section_path、section 等
doc_type: 文档类型 ('pdf' | 'word' | 'excel' | 'ppt')
Returns:
语义增强后的内容字符串,格式为:
标题
主题:章节路径
正文内容
"""
parts = []
title = getattr(chunk, 'title', '') or ''
if isinstance(title, list):
title = ' '.join(str(t) for t in title if t) or ''
if not isinstance(title, str):
title = ''
text_level = getattr(chunk, 'text_level', 0)
if title and title.strip() and text_level > 0:
parts.append(title.strip())
section = page_info.get('section_path', '') or page_info.get('section', '')
if isinstance(section, list):
section = ' > '.join(str(s) for s in section if s) or ''
if not isinstance(section, str):
section = ''
if section and section.strip():
section = _extract_section(section, max_levels=3)
parts.append(f"主题:{section}")
content = chunk.content if hasattr(chunk, 'content') else page_info.get('text', '')
if isinstance(content, list):
content = '\n'.join(str(item) for item in content)
parts.append(content)
return "\n".join(parts)
def _build_semantic_content_for_table(table_md: str, page_info: dict, chunk, doc_type: str) -> str:
"""
构建语义增强内容(表格类型)
将 Markdown 表格转换为包含语义摘要的增强格式,
提取表头、字段、行数和示例数据,提升向量检索命中率。
Args:
table_md: 表格的 Markdown 内容
page_info: 页面信息字典,含 section_path 等
chunk: 文档切片对象,含 title 属性
doc_type: 文档类型 ('pdf' | 'word' | 'excel' | 'ppt')
Returns:
语义增强后的内容字符串,包含:
主题:章节路径
表格:标题
字段:表头列表
描述:行数和字段概要
示例:首行数据示例
"""
parts = []
section = page_info.get('section_path', '') or page_info.get('section', '')
if isinstance(section, list):
section = ' > '.join(str(s) for s in section if s) or ''
if not isinstance(section, str):
section = ''
if section and section.strip():
section = _extract_section(section, max_levels=3)
parts.append(f"主题:{section}")
caption = getattr(chunk, 'title', '') or ''
if isinstance(caption, list):
caption = ' '.join(str(c) for c in caption if c) or ''
if caption and caption.strip() and caption != "表格":
parts.append(f"表格:{caption.strip()}")
# 检测是否是 HTML 格式,如果是则转换为 Markdown
if '<table' in table_md.lower():
try:
from parsers.mineru_parser import html_table_to_markdown
table_md = html_table_to_markdown(table_md)
except Exception as e:
logger.warning(f"HTML 表格转换失败: {e}")
lines = table_md.split('\n')
headers = []
for line in lines:
if line.startswith('|') and '---' not in line:
headers = [h.strip() for h in line.split('|') if h.strip()]
if headers and len(headers) > 1:
parts.append(f"字段:{', '.join(headers)}")
break
row_count = len([l for l in lines if l.startswith('|') and '---' not in l])
if headers:
parts.append(f"描述:该表包含{row_count}行数据,记录各{', '.join(headers[:3])}信息")
else:
parts.append(f"描述:该表包含{row_count}行数据")
for line in lines:
if line.startswith('|') and '---' not in line and headers:
cells = [c.strip() for c in line.split('|') if c.strip()]
if cells and cells != headers:
example_parts = []
for i, h in enumerate(headers[:2]):
if i < len(cells):
example_parts.append(f"{h}={cells[i]}")
if example_parts:
parts.append(f"示例:{', '.join(example_parts)}")
break
# 添加完整表格内容Markdown 格式)
if lines and any(l.startswith('|') for l in lines):
parts.append("") # 空行分隔
parts.append("表格内容:")
parts.extend(lines)
return "\n".join(parts)
# ==================== BM25 索引管理 ====================
class BM25Index:
"""
BM25 关键词检索索引
基于 rank_bm25.BM25Okapi 实现,使用 jieba 进行中文分词。
支持文档的添加、检索、持久化和加载。
Attributes:
bm25: BM25Okapi 索引实例
ids: 文档 ID 列表
documents: 文档内容列表
metadatas: 文档元数据列表
Example:
>>> index = BM25Index()
>>> index.add_documents(
... ids=["doc1", "doc2"],
... documents=["财务报销流程", "请假审批制度"],
... metadatas=[{"source": "a.pdf"}, {"source": "b.pdf"}]
... )
>>> ids, docs, metas, scores = index.search("报销", top_k=5)
"""
def __init__(self) -> None:
"""初始化空的 BM25 索引"""
self.bm25: Optional[BM25Okapi] = None
self.ids: List[str] = []
self.documents: List[str] = []
self.metadatas: List[dict] = []
def tokenize(self, text: str) -> List[str]:
"""
使用 jieba 对文本进行分词
Args:
text: 待分词的文本
Returns:
分词结果列表
"""
return list(jieba.cut(text))
def add_documents(self, ids: List[str], documents: List[str], metadatas: List[dict]) -> None:
"""
添加文档到索引(会覆盖原有索引)
Args:
ids: 文档 ID 列表
documents: 文档内容列表
metadatas: 文档元数据列表
"""
self.ids = ids
self.documents = documents
self.metadatas = metadatas
if documents:
tokenized = [self.tokenize(doc) for doc in documents]
self.bm25 = BM25Okapi(tokenized)
def search(self, query: str, top_k: int = 10) -> Tuple[List[str], List[str], List[dict], List[float]]:
"""
检索与查询最相关的文档
Args:
query: 查询文本
top_k: 返回的最大文档数
Returns:
元组 (ids, documents, metadatas, scores)
- ids: 文档 ID 列表
- documents: 文档内容列表
- metadatas: 文档元数据列表
- scores: BM25 分数列表
"""
if not self.bm25 or not self.documents:
return [], [], [], []
tokenized_query = self.tokenize(query)
scores = self.bm25.get_scores(tokenized_query)
top_indices = np.argsort(scores)[::-1][:top_k]
return (
[self.ids[i] for i in top_indices],
[self.documents[i] for i in top_indices],
[self.metadatas[i] for i in top_indices],
[float(scores[i]) for i in top_indices]
)
def save(self, filepath: str) -> None:
"""
持久化索引到文件
Args:
filepath: 保存路径(.pkl 文件)
"""
data = {'ids': self.ids, 'documents': self.documents, 'metadatas': self.metadatas}
os.makedirs(os.path.dirname(filepath), exist_ok=True)
with open(filepath, 'wb') as f:
pickle.dump(data, f)
def load(self, filepath: str) -> bool:
"""
从文件加载索引
Args:
filepath: 索引文件路径(.pkl 文件)
Returns:
加载成功返回 True失败返回 False
"""
if not os.path.exists(filepath):
return False
try:
with open(filepath, 'rb') as f:
data = pickle.load(f)
self.ids = data.get('ids', [])
self.documents = data.get('documents', [])
self.metadatas = data.get('metadatas', [])
if self.documents:
tokenized = [self.tokenize(doc) for doc in self.documents]
self.bm25 = BM25Okapi(tokenized)
return True
except Exception as e:
logger.error(f"加载 BM25 索引失败: {e}")
return False
def clear(self) -> None:
"""清空索引数据"""
self.bm25 = None
self.ids = []
self.documents = []
self.metadatas = []